From 7c7bd619e9ddf0234a9afbee22dde8f268f9c8e3 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 3 Jun 2026 09:20:12 +0000 Subject: [PATCH 001/214] one pass --- .../common/kv_cache_mem_manager/__init__.py | 2 + .../common/kv_cache_mem_manager/allocator.py | 12 +- .../deepseek4_mem_manager.py | 203 +++++++++ lightllm/common/req_manager.py | 106 ++++- lightllm/models/__init__.py | 1 + lightllm/models/deepseek_v4/__init__.py | 0 lightllm/models/deepseek_v4/infer_struct.py | 27 ++ .../deepseek_v4/layer_infer/__init__.py | 0 .../deepseek_v4/layer_infer/attention.py | 55 +++ .../deepseek_v4/layer_infer/compressor.py | 156 +++++++ .../layer_infer/hyper_connection.py | 58 +++ .../layer_infer/post_layer_infer.py | 19 + .../layer_infer/pre_layer_infer.py | 22 + .../layer_infer/transformer_layer_infer.py | 270 ++++++++++++ .../deepseek_v4/layer_weights/__init__.py | 0 .../pre_and_post_layer_weight.py | 37 ++ .../layer_weights/transformer_layer_weight.py | 398 ++++++++++++++++++ lightllm/models/deepseek_v4/mem_manager.py | 12 + lightllm/models/deepseek_v4/model.py | 121 ++++++ .../deepseek_v4/triton_kernel/__init__.py | 0 .../triton_kernel/quant_convert.py | 93 ++++ .../deepseek_v4/triton_kernel/rotary_emb.py | 26 ++ .../server/router/model_infer/infer_batch.py | 2 + 23 files changed, 1614 insertions(+), 6 deletions(-) create mode 100644 lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py create mode 100644 lightllm/models/deepseek_v4/__init__.py create mode 100644 lightllm/models/deepseek_v4/infer_struct.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/__init__.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/attention.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/compressor.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/hyper_connection.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py create mode 100644 lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py create mode 100644 lightllm/models/deepseek_v4/layer_weights/__init__.py create mode 100644 lightllm/models/deepseek_v4/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py create mode 100644 lightllm/models/deepseek_v4/mem_manager.py create mode 100644 lightllm/models/deepseek_v4/model.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/__init__.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/quant_convert.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py diff --git a/lightllm/common/kv_cache_mem_manager/__init__.py b/lightllm/common/kv_cache_mem_manager/__init__.py index 05544e149a..95f7e8ab76 100644 --- a/lightllm/common/kv_cache_mem_manager/__init__.py +++ b/lightllm/common/kv_cache_mem_manager/__init__.py @@ -4,6 +4,7 @@ from .ppl_int4kv_mem_manager import PPLINT4KVMemoryManager from .deepseek2_mem_manager import Deepseek2MemoryManager from .deepseek3_2mem_manager import Deepseek3_2MemoryManager +from .deepseek4_mem_manager import DeepseekV4MemoryManager from .fp8_per_token_group_quant_deepseek3_2mem_manager import FP8PerTokenGroupQuantDeepseek3_2MemoryManager from .fp8_static_per_head_quant_mem_manager import FP8StaticPerHeadQuantMemManager from .fp8_static_per_tensor_quant_mem_manager import FP8StaticPerTensorQuantMemManager @@ -17,6 +18,7 @@ "PPLINT8KVMemoryManager", "Deepseek2MemoryManager", "Deepseek3_2MemoryManager", + "DeepseekV4MemoryManager", "FP8PerTokenGroupQuantDeepseek3_2MemoryManager", "FP8StaticPerHeadQuantMemManager", "FP8StaticPerTensorQuantMemManager", diff --git a/lightllm/common/kv_cache_mem_manager/allocator.py b/lightllm/common/kv_cache_mem_manager/allocator.py index 850c158778..0179ed2714 100644 --- a/lightllm/common/kv_cache_mem_manager/allocator.py +++ b/lightllm/common/kv_cache_mem_manager/allocator.py @@ -3,13 +3,13 @@ from lightllm.utils.dist_utils import get_current_rank_in_node from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger -from typing import Union, List +from typing import Union, List, Optional logger = init_logger(__name__) class KvCacheAllocator: - def __init__(self, size: int) -> None: + def __init__(self, size: int, shared_name: Optional[str] = None) -> None: self.size = size self.mem_state = torch.arange( 0, self.size, dtype=torch.int32, device="cpu", requires_grad=False, pin_memory=True @@ -26,9 +26,11 @@ def __init__(self, size: int) -> None: rank_in_node = get_current_rank_in_node() # 用共享内存进行共享,router 模块读取进行精确的调度估计, nccl port 作为一个单机中单实列的标记。防止冲突。 - self.shared_can_use_token_num = SharedInt( - f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}" - ) + # shared_name 为 None 时使用主 kv 池的默认名(router 调度据此估算);DeepSeek-V4 的压缩子池等 + # 需要各自独立的计数器,传入区别于主池的唯一名,避免多个 allocator 写同一个共享计数器。 + if shared_name is None: + shared_name = f"{get_unique_server_name()}_mem_manger_can_use_token_num_{rank_in_node}" + self.shared_can_use_token_num = SharedInt(shared_name) self.shared_can_use_token_num.set_value(self.can_use_mem_size) return diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py new file mode 100644 index 0000000000..900b551cec --- /dev/null +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -0,0 +1,203 @@ +import torch +from typing import Dict, List, Optional +from .deepseek2_mem_manager import Deepseek2MemoryManager +from .allocator import KvCacheAllocator +from lightllm.utils.dist_utils import get_current_rank_in_node +from lightllm.utils.envs_utils import get_unique_server_name +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +class _SubKvPool: + """DeepSeek-V4 压缩分支(c4 / c128)使用的轻量子池。 + + 一个独立的 KvCacheAllocator + 一块压缩 latent buffer,可选附带一块与 latent 1:1 的 + indexer-K buffer(仅 c4/CSA 层用)。刻意不继承 MemoryManager —— pd/shm/kv_move 等机制 + 对压缩池暂不需要,保持最小。布局与主 MLA latent 池一致(每槽多预留 1 行作 padding 哨兵)。 + """ + + def __init__( + self, + size: int, + dtype: torch.dtype, + head_num: int, + head_dim: int, + layer_num: int, + indexer_head_dim: int = 0, + shared_name: Optional[str] = None, + device: str = "cuda", + ): + self.size = size + self.dtype = dtype + self.head_num = head_num + self.head_dim = head_dim + self.layer_num = layer_num + self.indexer_head_dim = indexer_head_dim + + self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device=device) + if indexer_head_dim > 0: + self.index_k_buffer = torch.empty((layer_num, size + 1, indexer_head_dim), dtype=dtype, device=device) + else: + self.index_k_buffer = None + + self.allocator = KvCacheAllocator(size, shared_name=shared_name) + self.HOLD_TOKEN_MEMINDEX = size + + def alloc(self, need_size) -> torch.Tensor: + return self.allocator.alloc(need_size) + + def free(self, free_index) -> None: + self.allocator.free(free_index) + + def free_all(self) -> None: + self.allocator.free_all() + + def get_kv_buffer(self, layer_index: int) -> torch.Tensor: + return self.kv_buffer[layer_index] + + def get_index_k_buffer(self, layer_index: int) -> torch.Tensor: + assert self.index_k_buffer is not None, "this sub pool has no indexer-K buffer" + return self.index_k_buffer[layer_index] + + +class DeepseekV4MemoryManager(Deepseek2MemoryManager): + """DeepSeek-V4 KV 管理(锁定决策: SWA 全历史 + 不分页)。 + + - dense/SWA latent: 继承 Deepseek2 的单张量 MLA latent ``kv_buffer``(每 token 一槽,所有层 + 共享层轴,head_num==1)。SWA 分支靠 layer_infer 传 ``AttControl(use_sliding_window)`` + attn_sink + 读最近窗口;dense 槽为纯 latent,不挂 indexer-K(与 V3.2 区别)。 + - c4_pool / c128_pool: 两个独立 ``_SubKvPool``(window 粒度,1-token 分配)。c4 池附带 indexer-K。 + - 容量: 用闭式 ``get_cell_size()``(= 每个 dense token 在所有池上的总字节)让基类 ``profile_size`` + 直接得到 full_token = dense 池大小,再按 1/4、1/128 派生压缩池大小。 + - compressor 递归状态不在这里,放 DeepseekV4ReqManager(后续步骤)。 + """ + + # dense 写入沿用 Deepseek2MemOperator(拆 nope/rope);压缩写入算子随 layer_infer 一并补。 + # operator_class 继承自 Deepseek2MemoryManager(= Deepseek2MemOperator)。 + + def __init__( + self, + size, + dtype, + head_num, + head_dim, + layer_num, + compress_rates: List[int], + indexer_head_dim: int = 128, + always_copy=False, + mem_fraction=0.9, + ): + assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" + assert ( + len(compress_rates) == layer_num + ), f"compress_rates 长度 {len(compress_rates)} 必须等于 layer_num {layer_num}" + assert all(r in (0, 4, 128) for r in compress_rates), "compress_rates 取值只能是 0/4/128" + + self.compress_rates = list(compress_rates) + self.n_c4 = sum(1 for r in self.compress_rates if r == 4) + self.n_c128 = sum(1 for r in self.compress_rates if r == 128) + self.indexer_head_dim = indexer_head_dim + + # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) + self.layer_to_c4_idx: Dict[int, int] = {} + self.layer_to_c128_idx: Dict[int, int] = {} + c4 = c128 = 0 + for lid, r in enumerate(self.compress_rates): + if r == 4: + self.layer_to_c4_idx[lid] = c4 + c4 += 1 + elif r == 128: + self.layer_to_c128_idx[lid] = c128 + c128 += 1 + + super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) + + def get_cell_size(self): + # 返回“每个 dense(full) token 在所有池上的总字节”。基类 profile_size 用 + # size = available_bytes / get_cell_size(),于是直接得到 full_token = dense 池大小。 + elem = torch._utils._element_size(self.dtype) + latent_bytes = self.head_num * self.head_dim * elem # 每 token 每层 dense latent + dense = latent_bytes * self.layer_num # SWA 全历史: 所有层 + c4 = latent_bytes * self.n_c4 / 4 # c4 压缩 latent + c128 = latent_bytes * self.n_c128 / 128 # c128 压缩 latent + indexer = self.indexer_head_dim * elem * self.n_c4 / 4 # c4 indexer-K + return dense + c4 + c128 + indexer + + def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): + # dense/SWA latent(继承 Deepseek2: [layer_num, size+1, head_num, head_dim]) + super()._init_buffers(size, dtype, head_num, head_dim, layer_num) + self._init_compressed_pools(size, dtype, head_num, head_dim) + + def _init_compressed_pools(self, size, dtype, head_num, head_dim): + rank_in_node = get_current_rank_in_node() + server = get_unique_server_name() + + self.c4_size = (size + 4 - 1) // 4 + self.c128_size = (size + 128 - 1) // 128 + + self.c4_pool: Optional[_SubKvPool] = None + self.c128_pool: Optional[_SubKvPool] = None + if self.n_c4 > 0: + self.c4_pool = _SubKvPool( + size=self.c4_size, + dtype=dtype, + head_num=head_num, + head_dim=head_dim, + layer_num=self.n_c4, + indexer_head_dim=self.indexer_head_dim, + shared_name=f"{server}_dsv4_c4_can_use_token_num_{rank_in_node}", + ) + if self.n_c128 > 0: + self.c128_pool = _SubKvPool( + size=self.c128_size, + dtype=dtype, + head_num=head_num, + head_dim=head_dim, + layer_num=self.n_c128, + indexer_head_dim=0, + shared_name=f"{server}_dsv4_c128_can_use_token_num_{rank_in_node}", + ) + + logger.info( + f"DeepseekV4MemoryManager pools: dense={size} " + f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " + f"indexer_head_dim={self.indexer_head_dim}" + ) + + # dense latent 读取沿用父类 get_att_input_params。 + + def _pool_and_local_layer(self, layer_index: int): + r = self.compress_rates[layer_index] + if r == 4: + return self.c4_pool, self.layer_to_c4_idx[layer_index] + if r == 128: + return self.c128_pool, self.layer_to_c128_idx[layer_index] + raise AssertionError(f"layer {layer_index} (rate {r}) 不是压缩层,没有压缩池") + + def get_compressed_kv_buffer(self, layer_index: int) -> torch.Tensor: + pool, local_layer = self._pool_and_local_layer(layer_index) + return pool.get_kv_buffer(local_layer) + + def get_compressed_indexer_k_buffer(self, layer_index: int) -> torch.Tensor: + assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" + return self.c4_pool.get_index_k_buffer(self.layer_to_c4_idx[layer_index]) + + def alloc_c4(self, need_size) -> torch.Tensor: + return self.c4_pool.alloc(need_size) + + def alloc_c128(self, need_size) -> torch.Tensor: + return self.c128_pool.alloc(need_size) + + def free_c4(self, free_index) -> None: + self.c4_pool.free(free_index) + + def free_c128(self, free_index) -> None: + self.c128_pool.free(free_index) + + def free_all(self): + super().free_all() + if self.c4_pool is not None: + self.c4_pool.free_all() + if self.c128_pool is not None: + self.c128_pool.free_all() diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 01e9c4ad35..c8197401c1 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -3,7 +3,7 @@ from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig from lightllm.utils.log_utils import init_logger -from .kv_cache_mem_manager import MemoryManager +from .kv_cache_mem_manager import MemoryManager, DeepseekV4MemoryManager from typing import List, Optional, TYPE_CHECKING from lightllm.common.basemodel.triton_kernel.gen_sampling_params import token_id_counter from lightllm.common.basemodel.triton_kernel.gen_sampling_params import update_req_to_token_id_counter @@ -299,3 +299,107 @@ def copy_small_page_buffer_to_linear_att_state( self.req_to_conv_state.buffer[:, dest_req_idx, ...] = conv_state self.req_to_ssm_state.buffer[:, dest_req_idx, ...] = ssm_state return + + +class DeepseekV4ReqManager(ReqManager): + """DeepSeek-V4 的请求级管理(锁定决策: SWA 全历史 + 不分页)。 + + 在基类 ReqManager 之上补三类 V4 专有的 per-request 结构(均从 mem_manager 读取 n_c4/n_c128/ + layer_to_*_idx/head_dim 等,避免重复配置): + + * ``req_to_c4_indexs`` / ``req_to_c128_indexs`` —— (req, 窗口下标) -> 压缩池槽位。 + 窗口下标 = position // compress_rate;窗口关闭时由 layer-infer 写入,attention 读取前 + n_windows 列即该 req 的全部压缩条目槽。未填充列为 0(不会被读到,语义同 req_to_token_indexs)。 + * ``req_to_c4_state`` / ``req_to_c128_state`` / ``req_to_c4_indexer_state`` —— compressor 的 + “在途窗口”累加状态(per req、per 压缩层),fp32。形状为 + ``(kv_or_score, coff * ratio, coff * dim)``; c4 因 Ca/Cb overlap 取 ``coff=2``, + c128 取 ``coff=1``。score 初始化为 ``-inf``,与官方 reference compressor 的 + ``kv_state``/``score_state`` 对齐。 + * entry_count 不另存:= position // compress_rate,可由序列长度推出。 + """ + + def __init__(self, max_request_num, max_sequence_length, mem_manager: DeepseekV4MemoryManager): + super().__init__(max_request_num, max_sequence_length, mem_manager) + assert isinstance(mem_manager, DeepseekV4MemoryManager) + self.n_c4 = mem_manager.n_c4 + self.n_c128 = mem_manager.n_c128 + head_dim = mem_manager.head_dim + indexer_head_dim = mem_manager.indexer_head_dim + + # (req, 窗口) -> 压缩槽。列数取 ceil(max_seq / ratio) 留足余量。 + c4_windows = (max_sequence_length + 4 - 1) // 4 + c128_windows = (max_sequence_length + 128 - 1) // 128 + self.req_to_c4_indexs = torch.zeros((max_request_num + 1, c4_windows), dtype=torch.int32, device="cuda") + self.req_to_c128_indexs = torch.zeros((max_request_num + 1, c128_windows), dtype=torch.int32, device="cuda") + + # compressor 在途窗口累加状态(fp32): [kv_or_score, coff * ratio, coff * dim]. + state_dtype = torch.float32 + self.req_to_c4_state = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(2, 8, 2 * head_dim), + layer_num=self.n_c4, + device="cuda", + ) + self.req_to_c128_state = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(2, 128, head_dim), + layer_num=self.n_c128, + device="cuda", + ) + self.req_to_c4_indexer_state = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(2, 8, 2 * indexer_head_dim), + layer_num=self.n_c4, + device="cuda", + ) + self._init_all_score_state() + return + + def _init_all_score_state(self): + if self.n_c4 > 0: + self.req_to_c4_state.buffer[:, :, 1, ...].fill_(float("-inf")) + self.req_to_c4_indexer_state.buffer[:, :, 1, ...].fill_(float("-inf")) + if self.n_c128 > 0: + self.req_to_c128_state.buffer[:, :, 1, ...].fill_(float("-inf")) + return + + def _reset_compress_cache_req(self, cache: LayerCache, req_idx: int): + if cache.layer_num == 0: + return + cache.buffer[:, req_idx, 0, ...].fill_(0) + cache.buffer[:, req_idx, 1, ...].fill_(float("-inf")) + return + + def init_compress_state(self, req_idx: int): + """新请求开始时重置其 compressor 在途状态(对应 mamba 的 init_linear_att_state)。""" + if self.n_c4 > 0: + self._reset_compress_cache_req(self.req_to_c4_state, req_idx) + self._reset_compress_cache_req(self.req_to_c4_indexer_state, req_idx) + if self.n_c128 > 0: + self._reset_compress_cache_req(self.req_to_c128_state, req_idx) + return + + def get_c4_compress_state(self, layer_index: int) -> torch.Tensor: + local = self.mem_manager.layer_to_c4_idx[layer_index] + return self.req_to_c4_state.buffer[local] + + def get_c128_compress_state(self, layer_index: int) -> torch.Tensor: + local = self.mem_manager.layer_to_c128_idx[layer_index] + return self.req_to_c128_state.buffer[local] + + def get_c4_indexer_compress_state(self, layer_index: int) -> torch.Tensor: + local = self.mem_manager.layer_to_c4_idx[layer_index] + return self.req_to_c4_indexer_state.buffer[local] + + def free(self, free_req_indexes, free_token_index, free_c4_index=None, free_c128_index=None): + """释放 dense 槽(基类)+ 压缩槽。压缩槽由调用方(infer batch)从 req_to_c*_indexs 收集后传入, + 与基类用 free_token_index 传 dense 槽的方式一致。""" + super().free(free_req_indexes, free_token_index) + if free_c4_index is not None and len(free_c4_index) > 0: + self.mem_manager.free_c4(free_c4_index) + if free_c128_index is not None and len(free_c128_index) > 0: + self.mem_manager.free_c128(free_c128_index) + return diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index f619b1d88f..3d376d160d 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -20,6 +20,7 @@ from lightllm.models.phi3.model import Phi3TpPartModel from lightllm.models.deepseek2.model import Deepseek2TpPartModel from lightllm.models.deepseek3_2.model import Deepseek3_2TpPartModel +from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel from lightllm.models.glm4_moe_lite.model import Glm4MoeLiteTpPartModel from lightllm.models.internvl.model import ( InternVLLlamaTpPartModel, diff --git a/lightllm/models/deepseek_v4/__init__.py b/lightllm/models/deepseek_v4/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py new file mode 100644 index 0000000000..6bc402cd28 --- /dev/null +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -0,0 +1,27 @@ +import torch +from lightllm.common.basemodel import InferStateInfo + + +class DeepseekV4InferStateInfo(InferStateInfo): + """Per-token interleaved-rope cos/sin for the two rope variants (sliding / compressed), following + the gemma4 two-variant convention (_cos_cached_* -> position_cos_*). Also exposes the full compressed + cos/sin tables, which the KV compressor indexes at window positions (not per-token).""" + + def __init__(self): + super().__init__() + self.position_cos_sliding = None + self.position_sin_sliding = None + self.position_cos_compress = None + self.position_sin_compress = None + self.cos_compress_table = None + self.sin_compress_table = None + + def init_some_extra_state(self, model): + super().init_some_extra_state(model) # sets position_ids, b_q_seq_len, b_q_start_loc (prefill) + pos = self.position_ids + self.position_cos_sliding = torch.index_select(model._cos_cached_sliding, 0, pos) # [T, rope_dim//2] + self.position_sin_sliding = torch.index_select(model._sin_cached_sliding, 0, pos) + self.position_cos_compress = torch.index_select(model._cos_cached_compress, 0, pos) + self.position_sin_compress = torch.index_select(model._sin_cached_compress, 0, pos) + self.cos_compress_table = model._cos_cached_compress + self.sin_compress_table = model._sin_cached_compress diff --git a/lightllm/models/deepseek_v4/layer_infer/__init__.py b/lightllm/models/deepseek_v4/layer_infer/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4/layer_infer/attention.py b/lightllm/models/deepseek_v4/layer_infer/attention.py new file mode 100644 index 0000000000..a25a2aa3d1 --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/attention.py @@ -0,0 +1,55 @@ +import torch +import torch.nn.functional as F + +# DeepSeek-V4 attention: MLA with a single shared KV head (head_dim=512), per-head learnable attention +# sink, and a candidate set = sliding-window tokens (size `window`) ++ compressed KV entries. Pure-torch +# transcription of the bundled reference (inference/model.py Attention.forward + kernel.py sparse_attn). +# Correctness-first prefill path. head_dim=512 > 256 so FlashAttention is unusable anyway; a fused +# triton sparse-gather kernel is a perf follow-up. + + +def torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): + """Gather-then-softmax attention with a per-head sink, matching reference kernel.sparse_attn. + + q:[b,m,h,d], kv:[b,n,d] (single KV head shared over h), attn_sink:[h] (fp32), + topk_idxs:[b,m,K] int (-1 = invalid/skip). Returns o:[b,m,h,d]. + """ + b, m, h, d = q.shape + n = kv.shape[1] + K = topk_idxs.shape[-1] + idx = topk_idxs.clamp(min=0).long() # [b,m,K] + keys = torch.gather(kv.unsqueeze(1).expand(b, m, n, d), 2, idx.unsqueeze(-1).expand(b, m, K, d)) # [b,m,K,d] + qf, kf = q.float(), keys.float() + scores = torch.einsum("bmhd,bmkd->bmhk", qf, kf) * scale # [b,m,h,K] + valid = (topk_idxs != -1).unsqueeze(2) # [b,m,1,K] + scores = scores.masked_fill(~valid, float("-inf")) + mx = scores.amax(dim=-1, keepdim=True) # [b,m,h,1] + mx = torch.nan_to_num(mx, neginf=0.0) + ex = (scores - mx).exp() # [b,m,h,K] + denom = ex.sum(-1) + (attn_sink.view(1, 1, h) - mx.squeeze(-1)).exp() # [b,m,h] + o = torch.einsum("bmhk,bmkd->bmhd", ex, kf) / denom.unsqueeze(-1) + return o.to(q.dtype) + + +def build_prefill_topk_idxs(seqlen, window, ratio, n_window, device): + """Per-query candidate indices into [window_kv (n_window tokens) ++ compressed_kv (ncomp entries)]. + + Returns int32 [seqlen, window + ncomp] with -1 for invalid. Window part indexes the per-token KV + (here stored as tokens 0..seqlen-1, so n_window == seqlen); compressed part is offset by n_window. + For prompts where ncomp <= index_topk the indexer is a no-op, so all causally-valid compressed + entries are attended (matches the reference for short context). + """ + t = torch.arange(seqlen, device=device) + # sliding window: query t attends tokens [max(0, t-window+1) .. t] + j = torch.arange(n_window, device=device) + win = j.unsqueeze(0).expand(seqlen, n_window).clone() # [s, n_window] + win_valid = (j.unsqueeze(0) <= t.unsqueeze(1)) & (j.unsqueeze(0) > (t.unsqueeze(1) - window)) + win = torch.where(win_valid, win, torch.full_like(win, -1)) + if ratio: + ncomp = seqlen // ratio + c = torch.arange(ncomp, device=device) + comp_valid = c.unsqueeze(0) < ((t.unsqueeze(1) + 1) // ratio) # [s, ncomp] + comp_idx = (c.unsqueeze(0) + n_window).expand(seqlen, ncomp) + comp = torch.where(comp_valid, comp_idx, torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long)) + return torch.cat([win, comp], dim=1).int() + return win.int() diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py new file mode 100644 index 0000000000..902de113db --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -0,0 +1,156 @@ +import torch +import torch.nn.functional as F +from ..triton_kernel.rotary_emb import apply_rotary_emb + +# KV compressor: pools every `ratio` consecutive tokens into one compressed KV entry via gated +# (softmax) pooling + a learned absolute-position bias (ape), RMSNorm, and rope on the trailing +# rope_dim. ratio==4 uses overlapping windows (two-series Ca/Cb scheme). Pure-torch transcription of +# the bundled reference inference/model.py Compressor.forward for the prefill (start_pos==0) path. +# NOTE: the reference also applies an FP8/FP4 QAT activation sim to the compressed entry; omitted here +# for the correctness-first prefill path (negligible vs argmax; revisit if e2e diverges). + + +def _overlap_transform(tensor, ratio, d, value): + # tensor: [nwin, ratio, 2*d] -> [nwin, 2*ratio, d]; slots [ratio:]=Cb(current), [:ratio]=Ca(previous window) + nwin = tensor.shape[0] + out = tensor.new_full((nwin, 2 * ratio, d), value) + out[:, ratio:] = tensor[:, :, d:] + out[1:, :ratio] = tensor[:-1, :, :d] + return out + + +def _rmsnorm(x, weight, eps): + xf = x.float() + xf = xf * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) + return (xf * weight.float()).to(x.dtype) + + +def compress_prefill(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, rope_dim, cos_table, sin_table, eps): + """x:[s,dim] (one request, start_pos=0) -> compressed kv [nwin, head_dim] (rope applied to last rope_dim). + + nwin = s // ratio (remainder tokens are decode-state, handled in the decode path). wkv_w/wgate_w: + [coff*head_dim, dim]; norm_w:[head_dim]; ape:[ratio, coff*head_dim]; cos_table/sin_table: compress rope tables. + """ + overlap = ratio == 4 + coff = 2 if overlap else 1 + d = head_dim + s = x.shape[0] + nwin = s // ratio + if nwin == 0: + # fewer than `ratio` tokens -> no completed window -> no compressed entry (matches reference) + return x.new_zeros(0, head_dim) + cutoff = nwin * ratio + xf = x.float() + kv = F.linear(xf, wkv_w.float())[:cutoff].view(nwin, ratio, coff * d) + score = F.linear(xf, wgate_w.float())[:cutoff].view(nwin, ratio, coff * d) + ape.float() + if overlap: + kv = _overlap_transform(kv, ratio, d, 0.0) + score = _overlap_transform(score, ratio, d, float("-inf")) + kv = (kv * torch.softmax(score, dim=1)).sum(dim=1) # [nwin, d] fp32 + kv = _rmsnorm(kv.to(x.dtype), norm_w, eps) # [nwin, d] + pos = torch.arange(nwin, device=x.device) * ratio + kv_rope = apply_rotary_emb(kv[:, -rope_dim:], cos_table[pos], sin_table[pos]) # cos/sin: [nwin, rope_dim//2] + return torch.cat([kv[:, :-rope_dim], kv_rope], dim=1) + + +def new_compressor_state(ratio, head_dim, device, dtype=torch.float32): + """Per-request compressor running state (matches reference Compressor.kv_state/score_state).""" + coff = 2 if ratio == 4 else 1 + kv_state = torch.zeros(coff * ratio, coff * head_dim, device=device, dtype=dtype) + score_state = torch.full((coff * ratio, coff * head_dim), float("-inf"), device=device, dtype=dtype) + return kv_state, score_state + + +def _finish_entry(kv, norm_w, ape_unused, rope_dim, cos_table, sin_table, position, eps, dtype): + kv = _rmsnorm(kv.to(dtype), norm_w, eps) # [d] + cos = cos_table[position : position + 1] # [1, rope_dim//2] + sin = sin_table[position : position + 1] + kv_rope = apply_rotary_emb(kv[-rope_dim:].unsqueeze(0), cos, sin)[0] + return torch.cat([kv[:-rope_dim], kv_rope], dim=0) + + +def compressor_prefill_state(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, rope_dim, cos_table, sin_table, eps): + """Faithful reference start_pos==0 path (incl. remainder). Returns (entries[ncomp,d], kv_state, score_state). + + entries have rope applied; kv_state/score_state carry the partial window for the decode path. + """ + overlap = ratio == 4 + coff = 2 if overlap else 1 + d = head_dim + s = x.shape[0] + dtype = x.dtype + xf = x.float() + kv = F.linear(xf, wkv_w.float()) # [s, coff*d] + score = F.linear(xf, wgate_w.float()) # [s, coff*d] + ape = ape.float() + kv_state, score_state = new_compressor_state(ratio, head_dim, x.device) + should_compress = s >= ratio + remainder = s % ratio + cutoff = s - remainder + offset = ratio if overlap else 0 + if overlap and cutoff >= ratio: + kv_state[:ratio] = kv[cutoff - ratio : cutoff] + score_state[:ratio] = score[cutoff - ratio : cutoff] + ape + if remainder > 0: + kv_state[offset : offset + remainder] = kv[cutoff:] + score_state[offset : offset + remainder] = score[cutoff:] + ape[:remainder] + kv = kv[:cutoff] + score = score[:cutoff] + if not should_compress: + return x.new_zeros(0, head_dim), kv_state, score_state + nwin = cutoff // ratio + kvw = kv.view(nwin, ratio, coff * d) + scw = score.view(nwin, ratio, coff * d) + ape + if overlap: + kvw = _overlap_transform(kvw, ratio, d, 0.0) + scw = _overlap_transform(scw, ratio, d, float("-inf")) + comp = (kvw * torch.softmax(scw, dim=1)).sum(dim=1) # [nwin, d] fp32 + comp = _rmsnorm(comp.to(dtype), norm_w, eps) + pos = torch.arange(nwin, device=x.device) * ratio + comp_rope = apply_rotary_emb(comp[:, -rope_dim:], cos_table[pos], sin_table[pos]) + comp = torch.cat([comp[:, :-rope_dim], comp_rope], dim=1) + return comp, kv_state, score_state + + +def compressor_decode_step( + x_new, + wkv_w, + wgate_w, + norm_w, + ape, + ratio, + head_dim, + rope_dim, + cos_table, + sin_table, + eps, + kv_state, + score_state, + start_pos, +): + """Faithful reference start_pos>0 path for one new token. Mutates kv_state/score_state in place. + Returns the new compressed entry [d] (rope applied) when a window completes, else None.""" + overlap = ratio == 4 + d = head_dim + dtype = x_new.dtype + xf = x_new.float().view(-1) # [dim] + kv = F.linear(xf, wkv_w.float()) # [coff*d] + score = F.linear(xf, wgate_w.float()) + ape.float()[start_pos % ratio] # [coff*d] + should_compress = (start_pos + 1) % ratio == 0 + if overlap: + kv_state[ratio + start_pos % ratio] = kv + score_state[ratio + start_pos % ratio] = score + if should_compress: + kv_cat = torch.cat([kv_state[:ratio, :d], kv_state[ratio:, d:]], dim=0) # [2*ratio, d] + sc_cat = torch.cat([score_state[:ratio, :d], score_state[ratio:, d:]], dim=0) + entry = (kv_cat * torch.softmax(sc_cat, dim=0)).sum(dim=0) # [d] + kv_state[:ratio] = kv_state[ratio:] + score_state[:ratio] = score_state[ratio:] + else: + kv_state[start_pos % ratio] = kv + score_state[start_pos % ratio] = score + if should_compress: + entry = (kv_state * torch.softmax(score_state, dim=0)).sum(dim=0) # [d] + if not should_compress: + return None + return _finish_entry(entry, norm_w, ape, rope_dim, cos_table, sin_table, start_pos + 1 - ratio, eps, dtype) diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py new file mode 100644 index 0000000000..75f540725b --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -0,0 +1,58 @@ +import torch +import torch.nn.functional as F + +# Manifold-constrained Hyper-Connections (mHC). Replaces the plain residual add: the hidden state is +# carried as ``hc_mult`` parallel streams. Each sub-layer (attn / ffn) collapses the streams to one +# vector (hc_pre), runs the sub-layer, then re-expands into the streams via learned post/comb weights +# (hc_post). A doubly-stochastic (Sinkhorn-normalized) ``comb`` matrix mixes the residual streams. +# Pure-torch transcription of the bundled reference inference/model.py (Block.hc_pre/hc_post, +# ParallelHead.hc_head) + inference/kernel.py (hc_split_sinkhorn). All math in fp32, as in the reference. + + +def hc_split_sinkhorn(mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps): + """mixes:[N, (2+hc)*hc] fp32 -> pre[N,hc], post[N,hc], comb[N,hc,hc] (doubly stochastic).""" + hc = hc_mult + pre = torch.sigmoid(mixes[:, :hc] * hc_scale[0] + hc_base[:hc]) + eps + post = 2.0 * torch.sigmoid(mixes[:, hc : 2 * hc] * hc_scale[1] + hc_base[hc : 2 * hc]) + comb = mixes[:, 2 * hc :].view(-1, hc, hc) * hc_scale[2] + hc_base[2 * hc :].view(hc, hc) + # comb = softmax(comb, dim=-1) + eps + comb = torch.softmax(comb, dim=-1) + eps + # one column normalization, then (iters-1) of (row, column) + comb = comb / (comb.sum(dim=-2, keepdim=True) + eps) + for _ in range(sinkhorn_iters - 1): + comb = comb / (comb.sum(dim=-1, keepdim=True) + eps) + comb = comb / (comb.sum(dim=-2, keepdim=True) + eps) + return pre, post, comb + + +def hc_pre(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps, sinkhorn_iters): + """streams:[N, hc*dim] -> (collapsed[N,dim], post[N,hc], comb[N,hc,hc]).""" + dtype = streams.dtype + x = streams.float() # [N, hc*dim] + rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + eps) + mixes = F.linear(x, hc_fn) * rsqrt # [N, (2+hc)*hc] + pre, post, comb = hc_split_sinkhorn(mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps) + streams3 = x.view(-1, hc_mult, dim) + collapsed = torch.sum(pre.unsqueeze(-1) * streams3, dim=1) # [N, dim] + return collapsed.to(dtype), post, comb + + +def hc_post(x, residual, post, comb, hc_mult, dim): + """x:[N,dim] sub-layer output, residual:[N, hc*dim] -> [N, hc*dim].""" + res = residual.float().view(-1, hc_mult, dim) # [N, hc, dim] + xf = x.float() + # post: [N,hc] -> [N,hc,dim]; comb mixes residual streams: out[i] = post[i]*x + sum_j comb[i,j]*res[j] + y = post.unsqueeze(-1) * xf.unsqueeze(-2) + torch.einsum("nij,njd->nid", comb, res) + return y.reshape(-1, hc_mult * dim).to(x.dtype) + + +def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps): + """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim] (sigmoid gate, no sinkhorn).""" + dtype = streams.dtype + x = streams.float() + rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + eps) + mixes = F.linear(x, hc_fn) * rsqrt # [N, hc] + pre = torch.sigmoid(mixes * hc_scale + hc_base) + eps # [N, hc] + streams3 = x.view(-1, hc_mult, dim) + collapsed = torch.sum(pre.unsqueeze(-1) * streams3, dim=1) + return collapsed.to(dtype) diff --git a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py new file mode 100644 index 0000000000..87951e7360 --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py @@ -0,0 +1,19 @@ +from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer +from .hyper_connection import hc_head + + +class DeepseekV4PostLayerInfer(LlamaPostLayerInfer): + """Collapse the hc_mult residual streams (hc_head) to [T, hidden], then final norm + lm_head.""" + + def token_forward(self, input_embdings, infer_state, layer_weight): + cfg = layer_weight.network_config_ + collapsed = hc_head( + input_embdings, + layer_weight.hc_head_fn_.weight, + layer_weight.hc_head_scale_.weight, + layer_weight.hc_head_base_.weight, + cfg["hc_mult"], + cfg["hidden_size"], + cfg.get("hc_eps", 1e-6), + ) + return super().token_forward(collapsed, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py new file mode 100644 index 0000000000..0be99ecbab --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py @@ -0,0 +1,22 @@ +import torch +import torch.distributed as dist +from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer +from lightllm.distributed.communication_op import all_reduce + + +class DeepseekV4PreLayerInfer(LlamaPreLayerInfer): + """Token embedding, then expand to the hc_mult parallel residual streams [T, hc_mult*hidden].""" + + def _embed_and_expand(self, input_ids, infer_state, layer_weight): + emb = layer_weight.wte_weight_(input_ids=input_ids, alloc_func=self.alloc_tensor) # [T, hidden] + if self.tp_world_size_ > 1: + all_reduce(emb, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + hc_mult = layer_weight.network_config_["hc_mult"] + t, hidden = emb.shape + return emb.unsqueeze(1).expand(t, hc_mult, hidden).reshape(t, hc_mult * hidden).contiguous() + + def context_forward(self, input_ids, infer_state, layer_weight): + return self._embed_and_expand(input_ids, infer_state, layer_weight) + + def token_forward(self, input_ids, infer_state, layer_weight): + return self._embed_and_expand(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py new file mode 100644 index 0000000000..a8dd0bb1e8 --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -0,0 +1,270 @@ +import torch +import torch.nn.functional as F +import torch.distributed as dist +from lightllm.common.basemodel import TransformerLayerInferTpl +from lightllm.distributed.communication_op import all_reduce +from lightllm.utils.envs_utils import get_env_start_args +from .hyper_connection import hc_pre, hc_post +from ..triton_kernel.rotary_emb import apply_rotary_emb +from .compressor import compressor_prefill_state, compressor_decode_step +from .attention import torch_sparse_attn +from ..triton_kernel.quant_convert import dequant_fp4_group_to_bf16 + + +class DeepseekV4TransformerLayerInfer(TransformerLayerInferTpl): + """One V4 decoder layer: HC(attn) then HC(ffn). Correctness-first pure-torch. + + The residual is carried as ``hc_mult`` streams flattened to [T, hc_mult*hidden]; each sub-layer + collapses (hc_pre), computes, and re-expands (hc_post). Attention is MLA over a sliding window + + compressed KV with a per-head sink (torch_sparse_attn); the MoE reuses lightllm's deepgemm FP8 + grouped GEMM driven by V4's custom router (sqrtsoftplus + hash/topk + bias-for-selection). + + Per-request decode state (window KV history + compressed KV + compressor running state) is kept in + a dict keyed by request id. NOTE: correctness-first — this should move into the KV mem manager for + production memory management / request eviction. + """ + + def __init__(self, layer_num, network_config): + super().__init__(layer_num, network_config) + cfg = network_config + self.eps_ = cfg["rms_norm_eps"] + self.hidden = cfg["hidden_size"] + self.n_heads = cfg["num_attention_heads"] + self.head_dim = cfg["head_dim"] + self.rope_dim = cfg["qk_rope_head_dim"] + self.o_groups = cfg["o_groups"] + self.o_lora = cfg["o_lora_rank"] + self.hc_mult = cfg["hc_mult"] + self.sinkhorn_iters = cfg["hc_sinkhorn_iters"] + self.hc_eps = cfg["hc_eps"] + self.window = cfg["sliding_window"] + self.compress_ratio = cfg["compress_ratios"][layer_num] + self.is_hash = layer_num < cfg["num_hash_layers"] + self.topk = cfg["num_experts_per_tok"] + self.route_scale = cfg["routed_scaling_factor"] + self.swiglu_limit = cfg["swiglu_limit"] + self.softmax_scale = self.head_dim**-0.5 + self.tp_q_heads = self.n_heads // self.tp_world_size_ + self.tp_groups = self.o_groups // self.tp_world_size_ + self.embed_dim_ = self.hc_mult * self.hidden + self.enable_ep_moe = get_env_start_args().enable_ep_moe + self._state = {} # req_id -> dict(kv_hist, comp_kv, cstate_kv, cstate_score) + + # ------------------------------------------------------------------ forward (HC-wrapped) + def _hc_block(self, streams, infer_state, lw, attn_fn): + residual = streams + collapsed, post, comb = hc_pre( + streams, + lw.hc_attn_fn_.weight, + lw.hc_attn_scale_.weight, + lw.hc_attn_base_.weight, + self.hc_mult, + self.hidden, + self.hc_eps, + self.sinkhorn_iters, + ) + o = attn_fn(lw.attn_norm_(collapsed, eps=self.eps_), infer_state, lw) + streams = hc_post(o, residual, post, comb, self.hc_mult, self.hidden) + + residual = streams + collapsed, post, comb = hc_pre( + streams, + lw.hc_ffn_fn_.weight, + lw.hc_ffn_scale_.weight, + lw.hc_ffn_base_.weight, + self.hc_mult, + self.hidden, + self.hc_eps, + self.sinkhorn_iters, + ) + f = self._moe_ffn(lw.ffn_norm_(collapsed, eps=self.eps_), infer_state, lw) + return hc_post(f, residual, post, comb, self.hc_mult, self.hidden) + + def context_forward(self, streams, infer_state, lw): + return self._hc_block(streams, infer_state, lw, self._attention_prefill) + + def token_forward(self, streams, infer_state, lw): + return self._hc_block(streams, infer_state, lw, self._attention_decode) + + # ------------------------------------------------------------------ shared projections + def _qkv(self, x, cos_tok, sin_tok, lw): + T = x.shape[0] + qa = lw.q_norm_(lw.wq_a_.mm(x), eps=self.eps_) + q = lw.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim).float() + q = (q * torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps_)).to(x.dtype) + q = torch.cat( + [ + q[..., : -self.rope_dim], + apply_rotary_emb(q[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1)), + ], + dim=-1, + ) + kv = lw.kv_norm_(lw.wkv_.mm(x), eps=self.eps_) + kv = torch.cat([kv[:, : -self.rope_dim], apply_rotary_emb(kv[:, -self.rope_dim :], cos_tok, sin_tok)], dim=1) + return q, kv + + def _out_proj(self, o, infer_state, lw): + # o: [T, tp_q_heads, head_dim] -> inverse rope -> grouped low-rank O -> [T, hidden] + T = o.shape[0] + o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] + o = lw.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] + o = lw.wo_b_.mm(o) + if self.tp_world_size_ > 1: + all_reduce(o, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + return o + + def _inv_rope(self, o, cos_tok, sin_tok): + return torch.cat( + [ + o[..., : -self.rope_dim], + apply_rotary_emb(o[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1), inverse=True), + ], + dim=-1, + ) + + # ------------------------------------------------------------------ attention (prefill) + def _attention_prefill(self, x, infer_state, lw): + T = x.shape[0] + if self.compress_ratio: + cos_tok, sin_tok = infer_state.position_cos_compress, infer_state.position_sin_compress + else: + cos_tok, sin_tok = infer_state.position_cos_sliding, infer_state.position_sin_sliding + q, kv = self._qkv(x, cos_tok, sin_tok, lw) + sink = lw.attn_sink_.weight + o = x.new_empty(T, self.tp_q_heads, self.head_dim) + b_req = infer_state.b_req_idx.tolist() + starts = infer_state.b_q_start_loc.tolist() + lens = infer_state.b_q_seq_len.tolist() + for req, st, ln in zip(b_req, starts, lens): + q_r, kv_r, x_r = q[st : st + ln], kv[st : st + ln], x[st : st + ln] + kv_all, n_window, ncomp = self._gather_prefill(x_r, kv_r, req, lw, infer_state) + ti = self._topk_idxs_prefill(ln, n_window, ncomp, x.device) + o[st : st + ln] = torch_sparse_attn(q_r.unsqueeze(0), kv_all.unsqueeze(0), sink, ti, self.softmax_scale)[0] + return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + + def _gather_prefill(self, x_r, kv_r, req, lw, infer_state): + ln = kv_r.shape[0] + if self.compress_ratio: + comp, ks, ss = compressor_prefill_state( + x_r, + lw.compressor_wkv_.mm_param.weight, + lw.compressor_wgate_.mm_param.weight, + lw.compressor_norm_.weight, + lw.compressor_ape_.weight, + self.compress_ratio, + self.head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + ) + self._state[req] = {"kv_hist": kv_r.detach(), "comp_kv": comp.detach(), "cstate_kv": ks, "cstate_score": ss} + return torch.cat([kv_r, comp], dim=0), ln, comp.shape[0] + self._state[req] = {"kv_hist": kv_r.detach()} + return kv_r, ln, 0 + + def _topk_idxs_prefill(self, seqlen, n_window, ncomp, device): + t = torch.arange(seqlen, device=device) + j = torch.arange(n_window, device=device) + win = torch.where( + (j.unsqueeze(0) <= t.unsqueeze(1)) & (j.unsqueeze(0) > (t.unsqueeze(1) - self.window)), + j.unsqueeze(0).expand(seqlen, n_window), + torch.full((seqlen, n_window), -1, device=device, dtype=torch.long), + ) + if ncomp: + c = torch.arange(ncomp, device=device) + comp = torch.where( + c.unsqueeze(0) < ((t.unsqueeze(1) + 1) // self.compress_ratio), + (c.unsqueeze(0) + n_window).expand(seqlen, ncomp), + torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long), + ) + return torch.cat([win, comp], dim=1).int().unsqueeze(0) + return win.int().unsqueeze(0) + + # ------------------------------------------------------------------ attention (decode) + def _attention_decode(self, x, infer_state, lw): + B = x.shape[0] # one new token per request + if self.compress_ratio: + cos_tok, sin_tok = infer_state.position_cos_compress, infer_state.position_sin_compress + else: + cos_tok, sin_tok = infer_state.position_cos_sliding, infer_state.position_sin_sliding + q, kv = self._qkv(x, cos_tok, sin_tok, lw) # [B, heads, hd], [B, hd] + sink = lw.attn_sink_.weight + b_req = infer_state.b_req_idx.tolist() + seqlens = infer_state.b_seq_len.tolist() + o = x.new_empty(B, self.tp_q_heads, self.head_dim) + for i, (req, seq) in enumerate(zip(b_req, seqlens)): + stt = self._state[req] + stt["kv_hist"] = torch.cat([stt["kv_hist"], kv[i : i + 1]], dim=0) + start_pos = seq - 1 + if self.compress_ratio: + e = compressor_decode_step( + x[i], + lw.compressor_wkv_.mm_param.weight, + lw.compressor_wgate_.mm_param.weight, + lw.compressor_norm_.weight, + lw.compressor_ape_.weight, + self.compress_ratio, + self.head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + stt["cstate_kv"], + stt["cstate_score"], + start_pos, + ) + if e is not None: + stt["comp_kv"] = torch.cat([stt["comp_kv"], e.unsqueeze(0)], dim=0) + win_kv = stt["kv_hist"][-self.window :] + kv_all = torch.cat([win_kv, stt["comp_kv"]], dim=0) + else: + win_kv = stt["kv_hist"][-self.window :] + kv_all = win_kv + ti = torch.arange(kv_all.shape[0], device=x.device).view(1, 1, -1).int() + o[i] = torch_sparse_attn( + q[i].view(1, 1, self.tp_q_heads, self.head_dim), kv_all.unsqueeze(0), sink, ti, self.softmax_scale + )[0, 0] + return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + + # ------------------------------------------------------------------ moe + def _fp4_experts(self, x, weights, indices, lw): + experts = lw.experts_ + out = torch.zeros(x.shape, device=x.device, dtype=torch.float32) + counts = torch.bincount(indices.reshape(-1), minlength=experts.n_routed_experts) + for expert_id in torch.nonzero(counts, as_tuple=False).flatten().tolist(): + token_idx, top_idx = torch.where(indices == expert_id) + if token_idx.numel() == 0: + continue + x_i = x[token_idx] + w1 = dequant_fp4_group_to_bf16(experts.w1[expert_id], experts.w1_scale[expert_id]) + w3 = dequant_fp4_group_to_bf16(experts.w3[expert_id], experts.w3_scale[expert_id]) + gate = F.linear(x_i, w1).float().clamp(max=self.swiglu_limit) + up = F.linear(x_i, w3).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) + hidden = F.silu(gate) * up + hidden.mul_(weights[token_idx, top_idx].unsqueeze(-1)) + w2 = dequant_fp4_group_to_bf16(experts.w2[expert_id], experts.w2_scale[expert_id]) + out.index_add_(0, token_idx, F.linear(hidden.to(x.dtype), w2).float()) + return out.to(x.dtype) + + def _moe_ffn(self, x, infer_state, lw): + gw = lw.gate_weight_.mm_param.weight + scores = F.softplus(F.linear(x.float(), gw.float())).sqrt() # sqrtsoftplus + if self.is_hash: + indices = lw.gate_tid2eid_.weight[infer_state.input_ids.long()] + else: + indices = (scores + lw.gate_bias_.weight.unsqueeze(0)).topk(self.topk, dim=-1)[1] + weights = scores.gather(1, indices) + weights = (weights / (weights.sum(-1, keepdim=True) + 1e-20) * self.route_scale).to(torch.float32) + routed = self._fp4_experts(x, weights, indices.long(), lw) + g = lw.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) + u = lw.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) + shared = lw.shared_down_.mm((F.silu(g) * u).to(x.dtype)) + if self.enable_ep_moe and getattr(lw.experts_, "is_ep", False): + if self.tp_world_size_ > 1: + all_reduce(shared, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + return routed + shared + out = routed + shared + if self.tp_world_size_ > 1: + all_reduce(out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + return out diff --git a/lightllm/models/deepseek_v4/layer_weights/__init__.py b/lightllm/models/deepseek_v4/layer_weights/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..54f29ce574 --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,37 @@ +import torch +from lightllm.common.basemodel import PreAndPostLayerWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + EmbeddingWeight, + LMHeadWeight, + RMSNormWeight, + ParameterWeight, +) + + +class DeepseekV4PreAndPostLayerWeight(PreAndPostLayerWeight): + def __init__(self, data_type, network_config): + super().__init__(data_type, network_config) + + hidden = network_config["hidden_size"] + vocab = network_config["vocab_size"] + hc_mult = network_config["hc_mult"] + + # embeddings / lm_head / final norm (bf16, vocab tensor-parallel). V4 has no `model.` prefix + # and does not tie embeddings (tie_word_embeddings=false). + self.wte_weight_ = EmbeddingWeight( + dim=hidden, vocab_size=vocab, weight_name="embed.weight", data_type=self.data_type_ + ) + self.lm_head_weight_ = LMHeadWeight( + dim=hidden, vocab_size=vocab, weight_name="head.weight", data_type=self.data_type_ + ) + self.final_norm_weight_ = RMSNormWeight(dim=hidden, weight_name="norm.weight", data_type=self.data_type_) + + # final hyper-connection head (collapses the hc_mult residual streams before the lm_head) + self.hc_head_fn_ = ParameterWeight( + weight_name="hc_head_fn", data_type=torch.float32, weight_shape=(hc_mult, hc_mult * hidden) + ) + self.hc_head_base_ = ParameterWeight( + weight_name="hc_head_base", data_type=torch.float32, weight_shape=(hc_mult,) + ) + self.hc_head_scale_ = ParameterWeight(weight_name="hc_head_scale", data_type=torch.float32, weight_shape=(1,)) + return diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py new file mode 100644 index 0000000000..7c12f714db --- /dev/null +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -0,0 +1,398 @@ +import torch +from lightllm.common.basemodel import TransformerLayerWeight +from lightllm.common.basemodel.layer_weights.meta_weights import ( + ROWMMWeight, + COLMMWeight, + ROWBMMWeight, + RMSNormWeight, + ParameterWeight, + TpAttSinkWeight, +) +from lightllm.common.basemodel.layer_weights.meta_weights.base_weight import BaseWeightTpl +from lightllm.common.quantization.registry import QUANTMETHODS +from ..triton_kernel.quant_convert import dequant_fp8_block_to_bf16 + + +class DeepseekV4FP4ExpertsWeight(BaseWeightTpl): + def __init__(self, weight_prefix, n_routed_experts, hidden_size, moe_intermediate_size, data_type): + super().__init__(data_type=data_type) + self.weight_prefix = weight_prefix + self.n_routed_experts = n_routed_experts + self.hidden_size = hidden_size + self.moe_intermediate_size = moe_intermediate_size + self.split_inter_size = moe_intermediate_size // self.tp_world_size_ + self.local_expert_ids = list(range(n_routed_experts)) + self.expert_idx_to_local_idx = {expert_idx: expert_idx for expert_idx in self.local_expert_ids} + self._create_weight() + + def _create_weight(self): + device = f"cuda:{self.device_id_}" + n = self.n_routed_experts + h = self.hidden_size + inter = self.split_inter_size + self.w1 = torch.empty((n, inter, h // 2), dtype=torch.int8, device=device) + self.w3 = torch.empty((n, inter, h // 2), dtype=torch.int8, device=device) + self.w2 = torch.empty((n, h, inter // 2), dtype=torch.int8, device=device) + self.w1_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) + self.w3_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) + self.w2_scale = torch.empty((n, h, inter // 32), dtype=torch.float8_e8m0fnu, device=device) + self.load_ok = { + name: [False] * n + for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale") + } + + def _copy_expert_weight(self, dst, weight, expert_idx, name, is_down=False): + if is_down: + start = self.tp_rank_ * self.split_inter_size // 2 + end = (self.tp_rank_ + 1) * self.split_inter_size // 2 + src = weight[:, start:end] + else: + start = self.tp_rank_ * self.split_inter_size + end = (self.tp_rank_ + 1) * self.split_inter_size + src = weight[start:end, :] + dst[expert_idx].copy_(src) + self.load_ok[name][expert_idx] = True + + def _copy_expert_scale(self, dst, scale, expert_idx, name, is_down=False): + if is_down: + start = self.tp_rank_ * self.split_inter_size // 32 + end = (self.tp_rank_ + 1) * self.split_inter_size // 32 + src = scale[:, start:end] + else: + start = self.tp_rank_ * self.split_inter_size + end = (self.tp_rank_ + 1) * self.split_inter_size + src = scale[start:end, :] + dst[expert_idx].copy_(src) + self.load_ok[name][expert_idx] = True + + def load_hf_weights(self, weights): + for expert_idx in self.local_expert_ids: + prefix = f"{self.weight_prefix}.{expert_idx}" + w1 = f"{prefix}.w1.weight" + w1_scale = f"{prefix}.w1.scale" + w2 = f"{prefix}.w2.weight" + w2_scale = f"{prefix}.w2.scale" + w3 = f"{prefix}.w3.weight" + w3_scale = f"{prefix}.w3.scale" + if w1 in weights: + self._copy_expert_weight(self.w1, weights[w1], expert_idx, "w1") + if w1_scale in weights: + self._copy_expert_scale(self.w1_scale, weights[w1_scale], expert_idx, "w1_scale") + if w3 in weights: + self._copy_expert_weight(self.w3, weights[w3], expert_idx, "w3") + if w3_scale in weights: + self._copy_expert_scale(self.w3_scale, weights[w3_scale], expert_idx, "w3_scale") + if w2 in weights: + self._copy_expert_weight(self.w2, weights[w2], expert_idx, "w2", is_down=True) + if w2_scale in weights: + self._copy_expert_scale(self.w2_scale, weights[w2_scale], expert_idx, "w2_scale", is_down=True) + + def verify_load(self): + return all(all(ok_list) for ok_list in self.load_ok.values()) + + +class DeepseekV4TransformerLayerWeight(TransformerLayerWeight): + """Per-layer weights for DeepSeek-V4-Flash. + + The checkpoint stores most linears in FP8 (e4m3 + block-128 ue8m0 scale) and the routed + experts in FP4 (int8-packed e2m1 + group-32 ue8m0 scale). Hopper does not use the SM100 + MegaMoE path here, so routed experts are kept in packed FP4 and temporarily de-quantized only + for selected experts in the correctness-first torch MoE path. + """ + + def __init__(self, layer_num, data_type, network_config, quant_cfg=None): + super().__init__(layer_num, data_type, network_config, quant_cfg) + return + + def _parse_config(self): + cfg = self.network_config_ + self.fp8_quant = QUANTMETHODS.get("deepgemm-fp8w8a8-b128") + self.hidden = cfg["hidden_size"] + self.n_heads = cfg["num_attention_heads"] + self.head_dim = cfg["head_dim"] + self.rope_dim = cfg["qk_rope_head_dim"] + self.q_lora_rank = cfg["q_lora_rank"] + self.o_lora_rank = cfg["o_lora_rank"] + self.o_groups = cfg["o_groups"] + self.index_n_heads = cfg["index_n_heads"] + self.index_head_dim = cfg["index_head_dim"] + self.n_routed_experts = cfg["n_routed_experts"] + self.moe_inter = cfg["moe_intermediate_size"] + self.num_hash_layers = cfg["num_hash_layers"] + self.vocab_size = cfg["vocab_size"] + self.hc_mult = cfg["hc_mult"] + self.mix_hc = (2 + self.hc_mult) * self.hc_mult + self.compress_ratio = cfg["compress_ratios"][self.layer_num_] + self.has_compressor = self.compress_ratio != 0 + self.has_indexer = self.compress_ratio == 4 + self.is_hash = self.layer_num_ < self.num_hash_layers + assert self.n_heads % self.tp_world_size_ == 0 + assert self.o_groups % self.tp_world_size_ == 0 + assert self.index_n_heads % self.tp_world_size_ == 0 + self.prefix = f"layers.{self.layer_num_}" + + def _init_weight_names(self): + return + + def _init_weight(self): + self._init_attn() + if self.has_compressor: + self._init_compressor(f"{self.prefix}.attn.compressor", self.head_dim, self.compress_ratio) + if self.has_indexer: + self._init_indexer() + self._init_moe() + self._init_norm() + self._init_hyper_connection() + + # ------------------------------------------------------------------ attention + def _init_attn(self): + p = f"{self.prefix}.attn" + # q low-rank (a replicated, b column-parallel over heads), kv single head (replicated) + self.wq_a_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.q_lora_rank], + weight_names=f"{p}.wq_a.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + tp_rank=0, + tp_world_size=1, + ) + self.wq_b_ = ROWMMWeight( + in_dim=self.q_lora_rank, + out_dims=[self.n_heads * self.head_dim], + weight_names=f"{p}.wq_b.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + self.wkv_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.head_dim], + weight_names=f"{p}.wkv.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + tp_rank=0, + tp_world_size=1, + ) + self.q_norm_ = RMSNormWeight(dim=self.q_lora_rank, weight_name=f"{p}.q_norm.weight", data_type=self.data_type_) + self.kv_norm_ = RMSNormWeight(dim=self.head_dim, weight_name=f"{p}.kv_norm.weight", data_type=self.data_type_) + self.attn_sink_ = TpAttSinkWeight( + all_q_head_num=self.n_heads, weight_name=f"{p}.attn_sink", data_type=torch.float32 + ) + # grouped low-rank output projection: wo_a is a per-group batched matmul [groups, in, o_lora], + # wo_b is row-parallel [groups*o_lora -> hidden]. wo_a is reshaped in load_hf_weights. + per_group_in = self.n_heads * self.head_dim // self.o_groups + self.wo_a_ = ROWBMMWeight( + dim0=self.o_groups, + dim1=per_group_in, + dim2=self.o_lora_rank, + weight_names=f"{p}.wo_a.weight", + data_type=self.data_type_, + quant_method=None, + ) + self.wo_b_ = COLMMWeight( + in_dim=self.o_groups * self.o_lora_rank, + out_dims=[self.hidden], + weight_names=f"{p}.wo_b.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + + # ------------------------------------------------------------------ compressor / indexer + def _init_compressor(self, prefix, head_dim, ratio): + coff = 2 if ratio == 4 else 1 + # wkv/wgate are bf16 (no scale) and replicated (single KV head). + self.compressor_wkv_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[coff * head_dim], + weight_names=f"{prefix}.wkv.weight", + data_type=self.data_type_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + self.compressor_wgate_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[coff * head_dim], + weight_names=f"{prefix}.wgate.weight", + data_type=self.data_type_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + self.compressor_norm_ = RMSNormWeight( + dim=head_dim, weight_name=f"{prefix}.norm.weight", data_type=self.data_type_ + ) + self.compressor_ape_ = ParameterWeight( + weight_name=f"{prefix}.ape", data_type=torch.float32, weight_shape=(ratio, coff * head_dim) + ) + + def _init_indexer(self): + p = f"{self.prefix}.attn.indexer" + # wq_b is FP8 in the checkpoint -> de-quantized to bf16 at load; column-parallel over index heads. + self.idx_wq_b_ = ROWMMWeight( + in_dim=self.q_lora_rank, + out_dims=[self.index_n_heads * self.index_head_dim], + weight_names=f"{p}.wq_b.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + self.idx_weights_proj_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.index_n_heads], + weight_names=f"{p}.weights_proj.weight", + data_type=self.data_type_, + quant_method=None, + ) + coff = 2 # indexer compressor always uses ratio 4 (overlap) + self.idx_cmp_wkv_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[coff * self.index_head_dim], + weight_names=f"{p}.compressor.wkv.weight", + data_type=self.data_type_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + self.idx_cmp_wgate_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[coff * self.index_head_dim], + weight_names=f"{p}.compressor.wgate.weight", + data_type=self.data_type_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + self.idx_cmp_norm_ = RMSNormWeight( + dim=self.index_head_dim, weight_name=f"{p}.compressor.norm.weight", data_type=self.data_type_ + ) + self.idx_cmp_ape_ = ParameterWeight( + weight_name=f"{p}.compressor.ape", data_type=torch.float32, weight_shape=(4, coff * self.index_head_dim) + ) + + # ------------------------------------------------------------------ moe + def _init_moe(self): + p = f"{self.prefix}.ffn" + # router gate (replicated) + self.gate_weight_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.n_routed_experts], + weight_names=f"{p}.gate.weight", + data_type=self.data_type_, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + if self.is_hash: + self.gate_tid2eid_ = ParameterWeight( + weight_name=f"{p}.gate.tid2eid", + data_type=torch.int64, + weight_shape=(self.vocab_size, self.network_config_["num_experts_per_tok"]), + ) + else: + self.gate_bias_ = ParameterWeight( + weight_name=f"{p}.gate.bias", data_type=torch.float32, weight_shape=(self.n_routed_experts,) + ) + # shared expert (dense, bf16 after de-quant): w1=gate, w3=up (row), w2=down (col) + sp = f"{p}.shared_experts" + self.shared_gate_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.moe_inter], + weight_names=f"{sp}.w1.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + self.shared_up_ = ROWMMWeight( + in_dim=self.hidden, + out_dims=[self.moe_inter], + weight_names=f"{sp}.w3.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + self.shared_down_ = COLMMWeight( + in_dim=self.moe_inter, + out_dims=[self.hidden], + weight_names=f"{sp}.w2.weight", + data_type=self.data_type_, + quant_method=self.fp8_quant, + ) + self.experts_ = DeepseekV4FP4ExpertsWeight( + weight_prefix=f"{p}.experts", + n_routed_experts=self.n_routed_experts, + hidden_size=self.hidden, + moe_intermediate_size=self.moe_inter, + data_type=self.data_type_, + ) + + def _init_norm(self): + self.attn_norm_ = RMSNormWeight( + dim=self.hidden, weight_name=f"{self.prefix}.attn_norm.weight", data_type=self.data_type_ + ) + self.ffn_norm_ = RMSNormWeight( + dim=self.hidden, weight_name=f"{self.prefix}.ffn_norm.weight", data_type=self.data_type_ + ) + + def _init_hyper_connection(self): + for which in ["attn", "ffn"]: + setattr( + self, + f"hc_{which}_fn_", + ParameterWeight( + weight_name=f"{self.prefix}.hc_{which}_fn", + data_type=torch.float32, + weight_shape=(self.mix_hc, self.hc_mult * self.hidden), + ), + ) + setattr( + self, + f"hc_{which}_base_", + ParameterWeight( + weight_name=f"{self.prefix}.hc_{which}_base", data_type=torch.float32, weight_shape=(self.mix_hc,) + ), + ) + setattr( + self, + f"hc_{which}_scale_", + ParameterWeight( + weight_name=f"{self.prefix}.hc_{which}_scale", data_type=torch.float32, weight_shape=(3,) + ), + ) + + # ------------------------------------------------------------------ loading + def load_hf_weights(self, weights): + self._dequant_in_place(weights) + return super().load_hf_weights(weights) + + def _direct_fp8_weight_names(self): + names = set() + for attr_name in dir(self): + attr = getattr(self, attr_name, None) + quant_method = getattr(attr, "quant_method", None) + if getattr(quant_method, "method_name", None) == "deepgemm-fp8w8a8-b128": + names.update(getattr(attr, "weight_names", [])) + return names + + def _dequant_in_place(self, weights): + p = self.prefix + "." + direct_fp8_names = self._direct_fp8_weight_names() + # Convert every (weight, scale) pair belonging to this layer. Existing FP8 matmul + # weights stay quantized; bmm-only weights are expanded; routed FP4 experts stay packed. + for k in [k for k in list(weights.keys()) if k.startswith(p) and k.endswith(".weight")]: + scale_k = k[: -len(".weight")] + ".scale" + if scale_k not in weights: + continue + w, s = weights[k], weights[scale_k] + if w.dtype == torch.int8: # FP4 routed experts stay packed for DeepseekV4FP4ExpertsWeight. + continue + elif k in direct_fp8_names: # FP8 e4m3, block-128 scale, run by DeepGEMM directly + weights[k.replace("weight", "weight_scale_inv")] = s.to(torch.float32) + del weights[scale_k] + else: # FP8 e4m3 for no-quant paths such as ROWBMMWeight + weights[k] = dequant_fp8_block_to_bf16(w, s).to(self.data_type_) + del weights[scale_k] + # grouped-O: reshape [groups*o_lora, in] -> [groups, in, o_lora] for the batched matmul + woa = f"{self.prefix}.attn.wo_a.weight" + if woa in weights and weights[woa].dim() == 2: + w = weights[woa] + per_group_in = self.n_heads * self.head_dim // self.o_groups + weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2).contiguous() + return diff --git a/lightllm/models/deepseek_v4/mem_manager.py b/lightllm/models/deepseek_v4/mem_manager.py new file mode 100644 index 0000000000..288d433380 --- /dev/null +++ b/lightllm/models/deepseek_v4/mem_manager.py @@ -0,0 +1,12 @@ +from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import Deepseek2MemoryManager + + +class DeepseekV4MemoryManager(Deepseek2MemoryManager): + """Stores the per-token MLA KV (head_num=1, head_dim=512), reusing the deepseek2 layout/operator. + + The prefill path computes attention in-layer from the request's hidden states, so it does not read + this buffer. The decode/incremental path (M6) will add the sliding-window ring + compressed-KV + + per-request compressor-state buffers here. + """ + + pass diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py new file mode 100644 index 0000000000..02c71f01b2 --- /dev/null +++ b/lightllm/models/deepseek_v4/model.py @@ -0,0 +1,121 @@ +import torch +from lightllm.models.registry import ModelRegistry +from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.common.req_manager import ReqManager, DeepseekV4ReqManager +from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager +from lightllm.common.basemodel.attention.triton.fp import TritonAttBackend +from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import DeepseekV4PreAndPostLayerWeight +from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight +from lightllm.models.deepseek_v4.layer_infer.pre_layer_infer import DeepseekV4PreLayerInfer +from lightllm.models.deepseek_v4.layer_infer.post_layer_infer import DeepseekV4PostLayerInfer +from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer +from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.llama.yarn_rotary_utils import find_correction_range, linear_ramp_mask +from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num +from lightllm.utils.log_utils import init_logger +from lightllm.distributed.communication_op import dist_group_manager + +logger = init_logger(__name__) + + +@ModelRegistry("deepseek_v4") +class DeepseekV4TpPartModel(LlamaTpPartModel): + + pre_and_post_weight_class = DeepseekV4PreAndPostLayerWeight + transformer_weight_class = DeepseekV4TransformerLayerWeight + + pre_layer_infer_class = DeepseekV4PreLayerInfer + post_layer_infer_class = DeepseekV4PostLayerInfer + transformer_layer_infer_class = DeepseekV4TransformerLayerInfer + + infer_state_class = DeepseekV4InferStateInfo + + def _verify_params(self): + assert self.load_way == "HF", "only support HF format weights" + assert self.config["num_attention_heads"] % self.tp_world_size_ == 0 + assert self.config["o_groups"] % self.tp_world_size_ == 0 + assert self.config["index_n_heads"] % self.tp_world_size_ == 0 + return + + def _init_some_value(self): + super()._init_some_value() + self.head_dim_ = self.config["head_dim"] + return + + def _init_req_manager(self): + create_max_seq_len = 0 + if self.batch_max_tokens is not None: + create_max_seq_len = max(create_max_seq_len, self.batch_max_tokens) + if self.max_seq_length is not None: + create_max_seq_len = max(create_max_seq_len, self.max_seq_length) + + self._dsv4_req_manager_seq_len = create_max_seq_len + self.req_manager = ReqManager(self.max_req_num, create_max_seq_len, None) + return + + def _get_compress_rates(self, layer_num): + rates = list(self.config.get("compress_ratios", [])) + if len(rates) < layer_num: + rates.extend([0] * (layer_num - len(rates))) + return rates[:layer_num] + + def _init_mem_manager(self): + layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() + self.mem_manager = DeepseekV4MemoryManager( + self.max_total_token_num, + dtype=self.data_type, + head_num=1, + head_dim=self.config["head_dim"], + layer_num=layer_num, + compress_rates=self._get_compress_rates(layer_num), + indexer_head_dim=self.config["index_head_dim"], + mem_fraction=self.mem_fraction, + ) + self.req_manager = DeepseekV4ReqManager( + self.max_req_num, self._dsv4_req_manager_seq_len, self.mem_manager + ) + return + + def _init_att_backend(self): + self.prefill_att_backend = TritonAttBackend(model=self) + self.decode_att_backend = TritonAttBackend(model=self) + return + + def _init_custom(self): + self._init_to_get_rotary() + dist_group_manager.new_deepep_group( + self.config["n_routed_experts"], + self.config["hidden_size"], + self.config.get("num_experts_per_tok", 1), + self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + ) + return + + def _init_to_get_rotary(self): + # Interleaved (GPT-J) rope. Build real cos/sin tables (_cos_cached_*/_sin_cached_*) following the + # gemma4 two-variant convention; the infer-struct slices them into position_cos_*/position_sin_* + # and apply_rotary_emb (interleaved, NOT the NeoX rotary_emb_fwd) applies them. Sliding-window + # layers use base rope_theta (no YaRN); compressed (CSA/HCA) layers use compress_rope_theta with + # YaRN. Tables kept fp32 for accuracy (the apply upcasts anyway). + cfg = self.config + rs = cfg.get("rope_scaling", {}) or {} + dim = cfg["qk_rope_head_dim"] + beta_fast = rs.get("beta_fast", 32) + beta_slow = rs.get("beta_slow", 1) + max_seq = max(int(self.max_seq_length), int(cfg.get("max_position_embeddings", 8192))) + max_seq = min(max_seq, 1 << 18) # cap table size (256K) for correctness-first + + def build(base, factor, orig_max): + freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim)) + if orig_max > 0: + low, high = find_correction_range(beta_fast, beta_slow, dim, base, orig_max) + smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() + freqs = freqs / factor * (1 - smooth) + freqs * smooth + f = torch.outer(torch.arange(max_seq, dtype=torch.float32, device="cuda"), freqs) # [max_seq, dim//2] + return f.cos(), f.sin() + + self._cos_cached_sliding, self._sin_cached_sliding = build(cfg["rope_theta"], rs.get("factor", 1.0), 0) + self._cos_cached_compress, self._sin_cached_compress = build( + cfg["compress_rope_theta"], rs.get("factor", 16), rs.get("original_max_position_embeddings", 65536) + ) + return diff --git a/lightllm/models/deepseek_v4/triton_kernel/__init__.py b/lightllm/models/deepseek_v4/triton_kernel/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py new file mode 100644 index 0000000000..c7d2d59ec6 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py @@ -0,0 +1,93 @@ +import torch + +# DeepSeek-V4-Flash ships weights in two quantized formats: +# * non-expert linears: FP8 e4m3 with block-[128,128] scales stored as float8_e8m0fnu (ue8m0) +# * routed experts: FP4 e2m1 packed 2-per-byte (stored as int8) with group-32 ue8m0 scales +# Hopper (H200) has no native SM100 MegaMoE path. Non-expert FP8 weights can run directly through +# DeepGEMM. Routed FP4 experts are converted blockwise to FP8, avoiding a full bf16 expansion. + +# OCP E2M1 magnitude table for the 3 low bits (sign = bit 3). torch.float4_e2m1fn_x2 packs two +# such codes per byte, low nibble = lower (even) logical index. +_E2M1_MAG = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0] + + +def e8m0_to_fp32(scale: torch.Tensor) -> torch.Tensor: + """float8_e8m0fnu encodes 2**(byte-127); torch decodes it correctly on .to(float32).""" + return scale.to(torch.float32) + + +def dequant_fp8_block_to_bf16(weight_e4m3: torch.Tensor, scale_e8m0: torch.Tensor, block_size: int = 128): + """De-quantize an FP8 e4m3 weight [out, in] with block-[bs,bs] ue8m0 scale to bf16.""" + from lightllm.models.deepseek2.triton_kernel.weight_dequant import weight_dequant + + w = weight_e4m3.cuda().contiguous() + s = e8m0_to_fp32(scale_e8m0).cuda().contiguous() + # weight_dequant runs with torch default dtype for the output; force bf16 result. + return weight_dequant(w, s, block_size) + + +def cast_e2m1fn_to_e4m3fn(weight_int8: torch.Tensor, scale_e8m0: torch.Tensor): + """Cast packed FP4 e2m1 expert weights to FP8 e4m3 with block-128 fp32 scales. + + This follows the DeepSeek-V4 reference converter, but returns the scale in fp32 because + LightLLM's DeepGEMM FP8 weight pack stores block scales as fp32. + """ + assert weight_int8.dtype == torch.int8 + assert weight_int8.ndim == 2 + out_dim, packed_in = weight_int8.shape + in_dim = packed_in * 2 + fp8_block_size = 128 + fp4_block_size = 32 + assert in_dim % fp8_block_size == 0 and out_dim % fp8_block_size == 0 + assert scale_e8m0.shape[0] == out_dim + assert scale_e8m0.shape[1] == in_dim // fp4_block_size + + table = torch.tensor( + [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0], + dtype=torch.float32, + device=weight_int8.device, + ) + packed = weight_int8.view(torch.uint8) + low = packed & 0x0F + high = (packed >> 4) & 0x0F + vals = torch.stack([table[low.long()], table[high.long()]], dim=-1).reshape(out_dim, in_dim) + + # 6.0 * 2**6 fits in e4m3fn (384 < 448), while 6.0 * 2**7 would overflow. + max_offset_bits = 6 + block_out = out_dim // fp8_block_size + block_in = in_dim // fp8_block_size + + vals = vals.view(block_out, fp8_block_size, block_in, fp8_block_size).transpose(1, 2) + scale = scale_e8m0.float().view(block_out, fp8_block_size, block_in, -1).transpose(1, 2).flatten(2) + block_scale = scale.amax(dim=-1, keepdim=True) / (2**max_offset_bits) + offset = scale / block_scale + offset = offset.unflatten(-1, (fp8_block_size, -1)).repeat_interleave(fp4_block_size, dim=-1) + vals = (vals * offset).transpose(1, 2).reshape(out_dim, in_dim) + block_scale = block_scale.squeeze(-1).to(torch.float8_e8m0fnu).to(torch.float32) + return vals.to(torch.float8_e4m3fn), block_scale + + +def dequant_fp4_group_to_bf16(weight_int8: torch.Tensor, scale_e8m0: torch.Tensor, group_size: int = 32): + """De-quantize an int8-packed FP4 e2m1 weight to bf16. + + weight_int8: [out, in // 2] int8 (two e2m1 codes per byte, low nibble = even index). + scale_e8m0: [out, in // group_size] ue8m0 (one scale per group_size logical elements along K). + returns: [out, in] bf16. + """ + w = weight_int8.cuda() + out, packed_in = w.shape + in_dim = packed_in * 2 + b = w.to(torch.int32).bitwise_and(0xFF) + lut = torch.tensor(_E2M1_MAG, dtype=torch.float32, device=w.device) + + def _decode(nib: torch.Tensor) -> torch.Tensor: + mag = lut[nib.bitwise_and(0x7)] + neg = nib.bitwise_and(0x8).bool() + return torch.where(neg, -mag, mag) + + lo = _decode(b.bitwise_and(0xF)) + hi = _decode(b.bitwise_right_shift(4).bitwise_and(0xF)) + vals = torch.stack([lo, hi], dim=-1).reshape(out, in_dim) # [out, in] + s = e8m0_to_fp32(scale_e8m0).cuda() # [out, in//group_size] + s = s.repeat_interleave(group_size, dim=1)[:, :in_dim] + return (vals * s).to(torch.bfloat16) diff --git a/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py b/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py new file mode 100644 index 0000000000..cb50977446 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py @@ -0,0 +1,26 @@ +import torch + +# Interleaved (GPT-J) rotary application for DeepSeek-V4. Unlike llama/gemma's NeoX-style +# rotary_emb_fwd (rotate-half: pairs channel i with i+d/2 over a real cos/sin table), V4 rotates +# adjacent pairs (x0,x1),(x2,x3),... — a different channel pairing — so it cannot reuse +# rotary_emb_fwd, but it consumes the same real cos/sin tables (built in model.py:_init_to_get_rotary +# as _cos_cached_*/_sin_cached_*, gemma4-style). Correctness-first pure-torch; a fused triton port is +# a perf follow-up. + + +def apply_rotary_emb(x, cos, sin, inverse=False): + """Apply interleaved rope to the LAST dim of x (size = 2*cos.size(-1)). + + x: [..., rope_dim] (real). cos/sin: [..., rope_dim//2], broadcastable to x's paired view. + For x of shape [N, H, rope_dim], pass cos/sin [N, 1, rope_dim//2]; for [N, rope_dim] pass [N, rope_dim//2]. + Returns a new tensor of x's dtype (not in-place). inverse=True applies the conjugate rotation. + """ + dtype = x.dtype + x = x.float().reshape(*x.shape[:-1], -1, 2) + x0, x1 = x[..., 0], x[..., 1] + cos = cos.float() + sin = sin.float() + if inverse: + sin = -sin + out = torch.stack([x0 * cos - x1 * sin, x0 * sin + x1 * cos], dim=-1) + return out.flatten(-2).to(dtype) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index f0ec69b2c1..5e90c9b34a 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -590,6 +590,8 @@ def _init_all_state(self): self.cur_output_len = 0 g_infer_context.req_manager.req_sampling_params_manager.init_req_sampling_params(self) + if hasattr(g_infer_context.req_manager, "init_compress_state"): + g_infer_context.req_manager.init_compress_state(req_idx=self.req_idx) self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() # token healing mode 才被使用的管理对象 From d790ad2407f75308fa30cf73f91dfbd928dd4306 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 5 Jun 2026 01:30:51 +0000 Subject: [PATCH 002/214] Optimization --- lightllm/__init__.py | 27 + .../deepseek4_mem_manager.py | 639 ++++++++++++++++-- .../kv_cache_mem_manager/operator/__init__.py | 1 + .../kv_cache_mem_manager/operator/deepseek.py | 19 +- lightllm/common/quantization/__init__.py | 3 + lightllm/common/req_manager.py | 279 +++++++- lightllm/models/deepseek3_2/model.py | 71 +- .../deepseek_v4/layer_infer/attention.py | 123 +++- .../deepseek_v4/layer_infer/compressor.py | 327 ++++++++- .../layer_infer/hyper_connection.py | 86 ++- .../layer_infer/transformer_layer_infer.py | 636 +++++++++++++++-- .../layer_weights/transformer_layer_weight.py | 155 ++++- lightllm/models/deepseek_v4/model.py | 144 +++- lightllm/server/api_start.py | 53 +- .../server/router/model_infer/infer_batch.py | 76 ++- lightllm/server/tokenizer.py | 5 + 16 files changed, 2295 insertions(+), 349 deletions(-) diff --git a/lightllm/__init__.py b/lightllm/__init__.py index e9ba6f3041..8e515afb70 100644 --- a/lightllm/__init__.py +++ b/lightllm/__init__.py @@ -1,4 +1,31 @@ from lightllm.utils.device_utils import is_musa + +def _patch_mp_resource_tracker_for_semaphore(): + from multiprocessing import resource_tracker + + if getattr(resource_tracker, "_lightllm_ignore_semaphore", False): + return + + orig_register = resource_tracker.register + orig_unregister = resource_tracker.unregister + + def register(name, rtype): + if rtype == "semaphore": + return + return orig_register(name, rtype) + + def unregister(name, rtype): + if rtype == "semaphore": + return + return orig_unregister(name, rtype) + + resource_tracker.register = register + resource_tracker.unregister = unregister + resource_tracker._lightllm_ignore_semaphore = True + + +_patch_mp_resource_tracker_for_semaphore() + if is_musa(): import torchada # noqa: F401 diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 900b551cec..739e0bd51a 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,44 +1,208 @@ import torch +import torch.distributed as dist from typing import Dict, List, Optional from .deepseek2_mem_manager import Deepseek2MemoryManager +from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator -from lightllm.utils.dist_utils import get_current_rank_in_node +from lightllm.utils.dist_utils import get_current_device_id, get_current_rank_in_node from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger +from lightllm.utils.profile_max_tokens import get_available_gpu_memory, get_total_gpu_memory logger = init_logger(__name__) -class _SubKvPool: - """DeepSeek-V4 压缩分支(c4 / c128)使用的轻量子池。 - - 一个独立的 KvCacheAllocator + 一块压缩 latent buffer,可选附带一块与 latent 1:1 的 - indexer-K buffer(仅 c4/CSA 层用)。刻意不继承 MemoryManager —— pd/shm/kv_move 等机制 - 对压缩池暂不需要,保持最小。布局与主 MLA latent 池一致(每槽多预留 1 行作 padding 哨兵)。 +DSV4_MLA_NOPE_DIM = 448 +DSV4_MLA_ROPE_DIM = 64 +DSV4_MLA_HEAD_DIM = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM +DSV4_MLA_QUANT_GROUP_SIZE = 64 +DSV4_MLA_SCALE_BYTES = DSV4_MLA_NOPE_DIM // DSV4_MLA_QUANT_GROUP_SIZE + 1 +DSV4_MLA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 + DSV4_MLA_SCALE_BYTES +DSV4_INDEXER_HEAD_DIM = 128 +DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + 4 +DSV4_FP8_E4M3_MAX = 448.0 +DSV4_FP8_SCALE_MIN = 1e-4 +DSV4_MLA_DATA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 +DSV4_MLA_SCALE_TAIL_BYTES = DSV4_MLA_SCALE_BYTES +DSV4_MLA_PAGE_ALIGN_BYTES = DSV4_MLA_DATA_BYTES_PER_TOKEN +DSV4_SWA_PAGE_SIZE = 128 +DSV4_C4_PAGE_SIZE = 64 +DSV4_C128_PAGE_SIZE = 2 +DSV4_PROFILE_MAX_FULL_TOKENS = 2_000_000 + + +def _ceil_div(a: int, b: int) -> int: + return (a + b - 1) // b + + +class _PageSlabMlaPool: + """SGLang-compatible fp8_ds_mla page-slab storage with token-slot addressing. + + The public loc is still a LightLLM token slot. Internally each page stores all + 576B NoPE+RoPE payloads first and the 8B scale records at the page tail: + data_offset = page * bytes_per_page + token_in_page * 576 + scale_offset = page * bytes_per_page + page_size * 576 + token_in_page * 8 """ def __init__( self, size: int, - dtype: torch.dtype, - head_num: int, - head_dim: int, + page_size: int, layer_num: int, - indexer_head_dim: int = 0, - shared_name: Optional[str] = None, device: str = "cuda", ): self.size = size - self.dtype = dtype - self.head_num = head_num - self.head_dim = head_dim + self.page_size = page_size self.layer_num = layer_num - self.indexer_head_dim = indexer_head_dim + self.dtype = torch.uint8 + self.data_bytes_per_token = DSV4_MLA_DATA_BYTES_PER_TOKEN + self.scale_bytes_per_token = DSV4_MLA_SCALE_TAIL_BYTES + self.bytes_per_token = DSV4_MLA_BYTES_PER_TOKEN + self.num_pages = _ceil_div(size + 1, page_size) + self.bytes_per_page = ( + _ceil_div(page_size * self.bytes_per_token, DSV4_MLA_PAGE_ALIGN_BYTES) * DSV4_MLA_PAGE_ALIGN_BYTES + ) + self.scale_offset_in_page = page_size * self.data_bytes_per_token + self.kv_buffer = torch.zeros( + (layer_num, self.num_pages, self.bytes_per_page), + dtype=torch.uint8, + device=device, + ) + self.HOLD_TOKEN_MEMINDEX = size + + def _loc_offsets(self, loc: torch.Tensor): + loc = loc.long() + page = torch.div(loc, self.page_size, rounding_mode="floor") + token = loc % self.page_size + page_base = page * self.bytes_per_page + data_offsets = page_base + token * self.data_bytes_per_token + scale_offsets = page_base + self.scale_offset_in_page + token * self.scale_bytes_per_token + return data_offsets, scale_offsets + + def write(self, layer_index: int, loc: torch.Tensor, packed: torch.Tensor) -> None: + if loc.numel() == 0: + return + loc = loc.long() + packed = packed.reshape(-1, DSV4_MLA_BYTES_PER_TOKEN).contiguous() + flat = self.kv_buffer[layer_index].view(-1) + data_offsets, scale_offsets = self._loc_offsets(loc) + + data = packed[:, : self.data_bytes_per_token].contiguous() + scale = packed[:, self.data_bytes_per_token : self.bytes_per_token].contiguous() + data_range = torch.arange(self.data_bytes_per_token, device=loc.device) + scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device) + flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)] = data + flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] = scale + return + + def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: + loc = loc.long() + if loc.numel() == 0: + return torch.empty((0, DSV4_MLA_BYTES_PER_TOKEN), dtype=torch.uint8, device=self.kv_buffer.device) + flat = self.kv_buffer[layer_index].view(-1) + data_offsets, scale_offsets = self._loc_offsets(loc) + data_range = torch.arange(self.data_bytes_per_token, device=loc.device) + scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device) + data = flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)] + scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] + return torch.cat([data, scale], dim=1).contiguous() + + def get_layer_buffer(self, layer_index: int) -> torch.Tensor: + return self.kv_buffer[layer_index] + - self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device=device) - if indexer_head_dim > 0: - self.index_k_buffer = torch.empty((layer_num, size + 1, indexer_head_dim), dtype=dtype, device=device) +class _PageSlabIndexerPool: + """C4 indexer-K storage: page tail stores per-token fp32 scales.""" + + def __init__( + self, + size: int, + page_size: int, + layer_num: int, + device: str = "cuda", + ): + self.size = size + self.page_size = page_size + self.layer_num = layer_num + self.head_dim = DSV4_INDEXER_HEAD_DIM + self.scale_bytes = 4 + self.bytes_per_token = DSV4_INDEXER_BYTES_PER_TOKEN + self.num_pages = _ceil_div(size + 1, page_size) + self.bytes_per_page = page_size * self.bytes_per_token + self.scale_offset_in_page = page_size * self.head_dim + self.index_k_buffer = torch.zeros( + (layer_num, self.num_pages, self.bytes_per_page), + dtype=torch.uint8, + device=device, + ) + self.HOLD_TOKEN_MEMINDEX = size + + def _loc_offsets(self, loc: torch.Tensor): + loc = loc.long() + page = torch.div(loc, self.page_size, rounding_mode="floor") + token = loc % self.page_size + page_base = page * self.bytes_per_page + k_offsets = page_base + token * self.head_dim + scale_offsets = page_base + self.scale_offset_in_page + token * self.scale_bytes + return k_offsets, scale_offsets + + def write(self, layer_index: int, loc: torch.Tensor, packed: torch.Tensor) -> None: + if loc.numel() == 0: + return + loc = loc.long() + packed = packed.reshape(-1, self.bytes_per_token).contiguous() + flat = self.index_k_buffer[layer_index].view(-1) + k_offsets, scale_offsets = self._loc_offsets(loc) + k_range = torch.arange(self.head_dim, device=loc.device) + scale_range = torch.arange(self.scale_bytes, device=loc.device) + flat[k_offsets.unsqueeze(1) + k_range.unsqueeze(0)] = packed[:, : self.head_dim] + flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] = packed[:, self.head_dim :] + return + + def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: + loc = loc.long() + if loc.numel() == 0: + return torch.empty((0, self.bytes_per_token), dtype=torch.uint8, device=self.index_k_buffer.device) + flat = self.index_k_buffer[layer_index].view(-1) + k_offsets, scale_offsets = self._loc_offsets(loc) + k_range = torch.arange(self.head_dim, device=loc.device) + scale_range = torch.arange(self.scale_bytes, device=loc.device) + k = flat[k_offsets.unsqueeze(1) + k_range.unsqueeze(0)] + scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] + return torch.cat([k, scale], dim=1).contiguous() + + def get_layer_buffer(self, layer_index: int) -> torch.Tensor: + return self.index_k_buffer[layer_index] + + +class _SubKvPool: + """Compressed c4/c128 KV pool with token-slot allocator and page-slab backing.""" + + def __init__( + self, + size: int, + page_size: int, + layer_num: int, + with_indexer: bool = False, + shared_name: Optional[str] = None, + device: str = "cuda", + ): + self.size = size + self.dtype = torch.uint8 + self.layer_num = layer_num + self.page_size = page_size + self.mla_pool = _PageSlabMlaPool(size=size, page_size=page_size, layer_num=layer_num, device=device) + self.kv_buffer = self.mla_pool.kv_buffer + if with_indexer: + self.indexer_pool = _PageSlabIndexerPool( + size=size, + page_size=page_size, + layer_num=layer_num, + device=device, + ) + self.index_k_buffer = self.indexer_pool.index_k_buffer else: + self.indexer_pool = None self.index_k_buffer = None self.allocator = KvCacheAllocator(size, shared_name=shared_name) @@ -54,27 +218,51 @@ def free_all(self) -> None: self.allocator.free_all() def get_kv_buffer(self, layer_index: int) -> torch.Tensor: - return self.kv_buffer[layer_index] + return self.mla_pool.get_layer_buffer(layer_index) def get_index_k_buffer(self, layer_index: int) -> torch.Tensor: - assert self.index_k_buffer is not None, "this sub pool has no indexer-K buffer" - return self.index_k_buffer[layer_index] + assert self.indexer_pool is not None, "this sub pool has no indexer-K buffer" + return self.indexer_pool.get_layer_buffer(layer_index) + + def write_kv(self, layer_index: int, slots: torch.Tensor, packed: torch.Tensor) -> None: + self.mla_pool.write(layer_index, slots, packed) + + def read_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + return self.mla_pool.read(layer_index, slots) + + def write_indexer_k(self, layer_index: int, slots: torch.Tensor, packed: torch.Tensor) -> None: + assert self.indexer_pool is not None + self.indexer_pool.write(layer_index, slots, packed) + + def read_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + assert self.indexer_pool is not None + return self.indexer_pool.read(layer_index, slots) class DeepseekV4MemoryManager(Deepseek2MemoryManager): - """DeepSeek-V4 KV 管理(锁定决策: SWA 全历史 + 不分页)。 - - - dense/SWA latent: 继承 Deepseek2 的单张量 MLA latent ``kv_buffer``(每 token 一槽,所有层 - 共享层轴,head_num==1)。SWA 分支靠 layer_infer 传 ``AttControl(use_sliding_window)`` + attn_sink - 读最近窗口;dense 槽为纯 latent,不挂 indexer-K(与 V3.2 区别)。 - - c4_pool / c128_pool: 两个独立 ``_SubKvPool``(window 粒度,1-token 分配)。c4 池附带 indexer-K。 - - 容量: 用闭式 ``get_cell_size()``(= 每个 dense token 在所有池上的总字节)让基类 ``profile_size`` - 直接得到 full_token = dense 池大小,再按 1/4、1/128 派生压缩池大小。 - - compressor 递归状态不在这里,放 DeepseekV4ReqManager(后续步骤)。 + """DeepSeek-V4 token-slot KV 管理(584B packed cache + bf16 workspace)。 + + - dense/SWA latent: 主 ``kv_buffer`` 仍是 LightLLM 的 token-slot cache,不分页;物理格式改为 + SGLang/vLLM 的 ``fp8_ds_mla``: 448B NoPE fp8 + 64*2B RoPE bf16 + 7B scale + 1B pad = 584B。 + - c4_pool / c128_pool: 两个独立 ``_SubKvPool``(window 粒度,1-token 分配),compressed KV 同样 + 存 584B packed。c4 池附带 132B/token 的 packed indexer-K。 + - 读取时先用 torch reference dequant/gather 回 bf16 workspace,供现有 vLLM sparse FlashMLA wrapper + 消费;下一步可把这些 pack/dequant helper 替换成 fused/triton 版本。 + - 容量: 用闭式 ``get_cell_size()``(= 每个 dense token 在所有池上的 packed 总字节)让基类 + ``profile_size`` 直接得到 full_token = dense 池大小,再按 1/4、1/128 派生压缩池大小。 + - compressor 递归状态放 DeepseekV4ReqManager。 """ - # dense 写入沿用 Deepseek2MemOperator(拆 nope/rope);压缩写入算子随 layer_infer 一并补。 - # operator_class 继承自 Deepseek2MemoryManager(= Deepseek2MemOperator)。 + operator_class = DeepseekV4MemOperator + + mla_nope_dim = DSV4_MLA_NOPE_DIM + mla_rope_dim = DSV4_MLA_ROPE_DIM + mla_head_dim = DSV4_MLA_HEAD_DIM + mla_quant_group_size = DSV4_MLA_QUANT_GROUP_SIZE + mla_scale_bytes = DSV4_MLA_SCALE_BYTES + mla_bytes_per_token = DSV4_MLA_BYTES_PER_TOKEN + indexer_head_dim_default = DSV4_INDEXER_HEAD_DIM + indexer_bytes_per_token = DSV4_INDEXER_BYTES_PER_TOKEN def __init__( self, @@ -85,19 +273,27 @@ def __init__( layer_num, compress_rates: List[int], indexer_head_dim: int = 128, + max_request_num: Optional[int] = None, + sliding_window: Optional[int] = None, always_copy=False, mem_fraction=0.9, ): assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" + assert head_dim == self.mla_head_dim, f"DeepSeek-V4 packed KV 期望 head_dim={self.mla_head_dim}" assert ( - len(compress_rates) == layer_num - ), f"compress_rates 长度 {len(compress_rates)} 必须等于 layer_num {layer_num}" + indexer_head_dim == self.indexer_head_dim_default + ), f"DeepSeek-V4 packed indexer-K 期望 indexer_head_dim={self.indexer_head_dim_default}" + assert len(compress_rates) == layer_num, f"compress_rates 长度 {len(compress_rates)} 必须等于 layer_num {layer_num}" assert all(r in (0, 4, 128) for r in compress_rates), "compress_rates 取值只能是 0/4/128" self.compress_rates = list(compress_rates) self.n_c4 = sum(1 for r in self.compress_rates if r == 4) self.n_c128 = sum(1 for r in self.compress_rates if r == 128) self.indexer_head_dim = indexer_head_dim + self.prefill_dtype = dtype + self.cache_dtype = torch.uint8 + self.max_request_num = max_request_num + self.sliding_window = sliding_window # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) self.layer_to_c4_idx: Dict[int, int] = {} @@ -113,23 +309,111 @@ def __init__( super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) + def _planned_swa_size(self, full_size: int) -> int: + if self.max_request_num is None or self.sliding_window is None: + return full_size + window_cap = max(1, int(self.max_request_num) * int(self.sliding_window)) + return max(1, min(full_size, window_cap)) + + def _dense_cell_size(self): + return self.head_num * self.mla_bytes_per_token * self.layer_num + + def _compressed_cell_size(self): + latent_bytes = self.head_num * self.mla_bytes_per_token + c4 = latent_bytes * self.n_c4 / 4 + c128 = latent_bytes * self.n_c128 / 128 + indexer = self.indexer_bytes_per_token * self.n_c4 / 4 + return c4 + c128 + indexer + + def profile_size(self, mem_fraction): + if self.size is not None: + return + + torch.cuda.empty_cache() + world_size = dist.get_world_size() + available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) + available_bytes = available_memory * 1024 ** 3 + dense_cell = self._dense_cell_size() + compressed_cell = self._compressed_cell_size() + + if self.max_request_num is not None and self.sliding_window is not None and compressed_cell > 0: + swa_cap = max(1, int(self.max_request_num) * int(self.sliding_window)) + full_cell = dense_cell + compressed_cell + bytes_until_swa_cap = full_cell * swa_cap + if available_bytes <= bytes_until_swa_cap: + self.size = max(1, int(available_bytes / full_cell)) + else: + self.size = max(1, int((available_bytes - dense_cell * swa_cap) / compressed_cell)) + else: + self.size = max(1, int(available_bytes / (dense_cell + compressed_cell))) + + if world_size > 1: + tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") + dist.all_reduce(tensor, op=dist.ReduceOp.MIN) + self.size = tensor.item() + + if self.size > DSV4_PROFILE_MAX_FULL_TOKENS: + logger.info( + f"DeepseekV4MemoryManager cap profiled max_total_token_num from " + f"{self.size} to {DSV4_PROFILE_MAX_FULL_TOKENS} to keep runtime headroom" + ) + self.size = DSV4_PROFILE_MAX_FULL_TOKENS + + logger.info( + f"{str(available_memory)} GB space is available after load the model weight\n" + f"{str((dense_cell + compressed_cell) / 1024 ** 2)} MB is the conservative size of one token kv cache\n" + f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" + ) + return + def get_cell_size(self): - # 返回“每个 dense(full) token 在所有池上的总字节”。基类 profile_size 用 - # size = available_bytes / get_cell_size(),于是直接得到 full_token = dense 池大小。 - elem = torch._utils._element_size(self.dtype) - latent_bytes = self.head_num * self.head_dim * elem # 每 token 每层 dense latent - dense = latent_bytes * self.layer_num # SWA 全历史: 所有层 - c4 = latent_bytes * self.n_c4 / 4 # c4 压缩 latent - c128 = latent_bytes * self.n_c128 / 128 # c128 压缩 latent - indexer = self.indexer_head_dim * elem * self.n_c4 / 4 # c4 indexer-K - return dense + c4 + c128 + indexer + dense = self._dense_cell_size() + compressed = self._compressed_cell_size() + if self.size is None: + return dense + compressed + swa_ratio = self._planned_swa_size(self.size) / max(1, self.size) + return dense * swa_ratio + compressed def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): - # dense/SWA latent(继承 Deepseek2: [layer_num, size+1, head_num, head_dim]) - super()._init_buffers(size, dtype, head_num, head_dim, layer_num) - self._init_compressed_pools(size, dtype, head_num, head_dim) + self.swa_size = self._planned_swa_size(size) + self.swa_pool = _PageSlabMlaPool( + size=self.swa_size, + page_size=DSV4_SWA_PAGE_SIZE, + layer_num=layer_num, + device="cuda", + ) + self.kv_buffer = self.swa_pool.kv_buffer + self._init_swa_mapping(size) + self._init_compressed_pools(size, head_num) + + def _init_swa_mapping(self, size): + rank_in_node = get_current_rank_in_node() + server = get_unique_server_name() + self.swa_allocator = KvCacheAllocator( + self.swa_size, + shared_name=f"{server}_dsv4_swa_can_use_token_num_{rank_in_node}", + ) + self.full_to_swa_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") + self.full_to_swa_indexs[size] = self.swa_pool.HOLD_TOKEN_MEMINDEX + if self.max_request_num is None or self.sliding_window is None: + self.req_to_swa_indexs = None + self.req_to_swa_full_indexs = None + return + + self.req_to_swa_indexs = torch.full( + (self.max_request_num + 1, self.sliding_window), + self.swa_pool.HOLD_TOKEN_MEMINDEX, + dtype=torch.int32, + device="cuda", + ) + self.req_to_swa_full_indexs = torch.full( + (self.max_request_num + 1, self.sliding_window), + -1, + dtype=torch.int32, + device="cuda", + ) - def _init_compressed_pools(self, size, dtype, head_num, head_dim): + def _init_compressed_pools(self, size, head_num): rank_in_node = get_current_rank_in_node() server = get_unique_server_name() @@ -141,31 +425,254 @@ def _init_compressed_pools(self, size, dtype, head_num, head_dim): if self.n_c4 > 0: self.c4_pool = _SubKvPool( size=self.c4_size, - dtype=dtype, - head_num=head_num, - head_dim=head_dim, + page_size=DSV4_C4_PAGE_SIZE, layer_num=self.n_c4, - indexer_head_dim=self.indexer_head_dim, + with_indexer=True, shared_name=f"{server}_dsv4_c4_can_use_token_num_{rank_in_node}", ) if self.n_c128 > 0: self.c128_pool = _SubKvPool( size=self.c128_size, - dtype=dtype, - head_num=head_num, - head_dim=head_dim, + page_size=DSV4_C128_PAGE_SIZE, layer_num=self.n_c128, - indexer_head_dim=0, + with_indexer=False, shared_name=f"{server}_dsv4_c128_can_use_token_num_{rank_in_node}", ) logger.info( - f"DeepseekV4MemoryManager pools: dense={size} " + f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size} " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " - f"indexer_head_dim={self.indexer_head_dim}" + f"packed_kv_bytes={self.mla_bytes_per_token} indexer_bytes={self.indexer_bytes_per_token}" ) - # dense latent 读取沿用父类 get_att_input_params。 + def get_att_input_params(self, layer_index: int): + return self.swa_pool.get_layer_buffer(layer_index) + + def _pack_mla_kv(self, kv: torch.Tensor) -> torch.Tensor: + kv = kv.reshape(-1, self.mla_head_dim) + out = torch.empty((kv.shape[0], self.mla_bytes_per_token), dtype=torch.uint8, device=kv.device) + nope = kv[:, : self.mla_nope_dim].float().reshape(-1, self.mla_scale_bytes - 1, self.mla_quant_group_size) + scale = torch.clamp(nope.abs().amax(dim=-1) / DSV4_FP8_E4M3_MAX, min=DSV4_FP8_SCALE_MIN) + scale_exp = torch.ceil(torch.log2(scale)).to(torch.int32) + scale = torch.exp2(scale_exp.float()) + nope_fp8 = torch.clamp(nope / scale.unsqueeze(-1), -DSV4_FP8_E4M3_MAX, DSV4_FP8_E4M3_MAX).to( + torch.float8_e4m3fn + ) + out[:, : self.mla_nope_dim].copy_(nope_fp8.reshape(-1, self.mla_nope_dim).view(dtype=torch.uint8)) + rope_start = self.mla_nope_dim + rope_end = rope_start + self.mla_rope_dim * 2 + rope = kv[:, self.mla_nope_dim : self.mla_head_dim].contiguous().to(torch.bfloat16) + out[:, rope_start:rope_end].copy_(rope.view(dtype=torch.uint8).reshape(-1, self.mla_rope_dim * 2)) + scale_start = rope_end + scale_end = scale_start + self.mla_scale_bytes - 1 + out[:, scale_start:scale_end].copy_((scale_exp + 127).to(torch.uint8)) + out[:, scale_end].zero_() + return out + + def _unpack_mla_kv(self, packed: torch.Tensor) -> torch.Tensor: + packed = packed.reshape(-1, self.mla_bytes_per_token) + if packed.shape[0] == 0: + return torch.empty((0, self.mla_head_dim), dtype=self.dtype, device=packed.device) + nope_fp8 = packed[:, : self.mla_nope_dim].view(dtype=torch.float8_e4m3fn).float() + nope_fp8 = nope_fp8.reshape(-1, self.mla_scale_bytes - 1, self.mla_quant_group_size) + rope_start = self.mla_nope_dim + rope_end = rope_start + self.mla_rope_dim * 2 + scale_start = rope_end + scale_end = scale_start + self.mla_scale_bytes - 1 + scale_exp = packed[:, scale_start:scale_end].to(torch.int32) - 127 + scale = torch.exp2(scale_exp.float()) + nope = (nope_fp8 * scale.reshape(-1, self.mla_scale_bytes - 1, 1)).reshape(-1, self.mla_nope_dim) + rope = packed[:, rope_start:rope_end].view(dtype=torch.bfloat16) + return torch.cat([nope.to(self.dtype), rope.to(self.dtype)], dim=-1) + + def _pack_indexer_k(self, indexer_k: torch.Tensor) -> torch.Tensor: + indexer_k = indexer_k.reshape(-1, self.indexer_head_dim) + out = torch.empty( + (indexer_k.shape[0], self.indexer_bytes_per_token), + dtype=torch.uint8, + device=indexer_k.device, + ) + k_float = indexer_k.float() + scale = torch.clamp( + k_float.abs().amax(dim=-1, keepdim=True) / DSV4_FP8_E4M3_MAX, + min=DSV4_FP8_SCALE_MIN, + ) + k_fp8 = torch.clamp(k_float / scale, -DSV4_FP8_E4M3_MAX, DSV4_FP8_E4M3_MAX).to(torch.float8_e4m3fn) + out[:, : self.indexer_head_dim].copy_(k_fp8.view(dtype=torch.uint8)) + out[:, self.indexer_head_dim : self.indexer_bytes_per_token].copy_(scale.view(dtype=torch.uint8).reshape(-1, 4)) + return out + + def _unpack_indexer_k(self, packed: torch.Tensor) -> torch.Tensor: + packed = packed.reshape(-1, self.indexer_bytes_per_token) + if packed.shape[0] == 0: + return torch.empty((0, self.indexer_head_dim), dtype=self.dtype, device=packed.device) + k_fp8 = packed[:, : self.indexer_head_dim].view(dtype=torch.float8_e4m3fn).float() + scale = packed[:, self.indexer_head_dim : self.indexer_bytes_per_token].view(dtype=torch.float32) + return (k_fp8 * scale).to(self.dtype) + + def _identity_swa_slots(self, full_slots: torch.Tensor) -> torch.Tensor: + full_slots = full_slots.long() + valid = full_slots != self.HOLD_TOKEN_MEMINDEX + if valid.any() and int(full_slots[valid].max().item()) >= self.swa_size: + raise RuntimeError( + "DeepSeek-V4 SWA cache needs req_idx/positions for full token slots outside the SWA pool" + ) + swa_slots = torch.where( + valid, + full_slots, + torch.full_like(full_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), + ) + if valid.any(): + self.full_to_swa_indexs[full_slots[valid]] = swa_slots[valid].to(torch.int32) + return swa_slots + + def ensure_swa_slots(self, req_idx: int, positions: torch.Tensor, full_slots: torch.Tensor) -> torch.Tensor: + full_slots = full_slots.long().reshape(-1) + if full_slots.numel() == 0: + return full_slots + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: + return self._identity_swa_slots(full_slots) + + positions = positions.long().reshape(-1) + assert positions.numel() == full_slots.numel() + req_idx = int(req_idx) + out = torch.empty_like(full_slots, dtype=torch.long) + for i, (pos, full) in enumerate(zip(positions.tolist(), full_slots.tolist())): + if full == self.HOLD_TOKEN_MEMINDEX: + out[i] = self.swa_pool.HOLD_TOKEN_MEMINDEX + continue + + ring_pos = pos % self.sliding_window + old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) + old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) + if old_full == full and old_swa != self.swa_pool.HOLD_TOKEN_MEMINDEX: + swa = old_swa + elif old_swa != self.swa_pool.HOLD_TOKEN_MEMINDEX: + if old_full >= 0: + self.full_to_swa_indexs[old_full] = -1 + swa = old_swa + else: + swa = int(self.swa_allocator.alloc(1)[0].item()) + + self.req_to_swa_indexs[req_idx, ring_pos] = swa + self.req_to_swa_full_indexs[req_idx, ring_pos] = full + self.full_to_swa_indexs[full] = swa + out[i] = swa + return out + + def _swa_slots_from_full(self, full_slots: torch.Tensor) -> torch.Tensor: + full_slots = full_slots.long().reshape(-1) + if full_slots.numel() == 0: + return full_slots + mapped = self.full_to_swa_indexs[full_slots].long() + missing = mapped < 0 + if missing.any(): + if self.req_to_swa_indexs is not None: + bad = int(full_slots[missing][0].item()) + raise RuntimeError(f"DeepSeek-V4 dense KV for full token slot {bad} has been evicted from SWA cache") + fallback = full_slots[missing] + fallback_valid = fallback < self.swa_size + if fallback_valid.all(): + mapped[missing] = fallback + self.full_to_swa_indexs[fallback] = fallback.to(torch.int32) + else: + bad = int(fallback[~fallback_valid][0].item()) + raise RuntimeError(f"DeepSeek-V4 dense KV for full token slot {bad} has been evicted from SWA cache") + return mapped + + def free_swa_for_req(self, req_idx: int) -> None: + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: + return + req_idx = int(req_idx) + slots = self.req_to_swa_indexs[req_idx] + full_slots = self.req_to_swa_full_indexs[req_idx] + valid_swa = slots != self.swa_pool.HOLD_TOKEN_MEMINDEX + if valid_swa.any(): + free_slots = torch.unique(slots[valid_swa]).detach().cpu() + self.swa_allocator.free(free_slots) + valid_full = full_slots >= 0 + if valid_full.any(): + self.full_to_swa_indexs[full_slots[valid_full].long()] = -1 + self.req_to_swa_indexs[req_idx].fill_(self.swa_pool.HOLD_TOKEN_MEMINDEX) + self.req_to_swa_full_indexs[req_idx].fill_(-1) + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + + def _keep_last_swa_writes(self, swa_slots: torch.Tensor, packed: torch.Tensor): + """Drop duplicate SWA writes generated by long prefill ring reuse.""" + if swa_slots.numel() <= 1: + return swa_slots, packed + + slots_cpu = swa_slots.detach().cpu().tolist() + seen = set() + keep = [] + hold = self.swa_pool.HOLD_TOKEN_MEMINDEX + for i in range(len(slots_cpu) - 1, -1, -1): + slot = int(slots_cpu[i]) + if slot == hold or slot in seen: + continue + seen.add(slot) + keep.append(i) + keep.reverse() + if len(keep) == len(slots_cpu): + return swa_slots, packed + if not keep: + return swa_slots[:0], packed[:0] + keep_index = torch.tensor(keep, dtype=torch.long, device=swa_slots.device) + return swa_slots.index_select(0, keep_index), packed.index_select(0, keep_index) + + def pack_mla_kv_to_cache( + self, + layer_index: int, + mem_index: torch.Tensor, + kv: torch.Tensor, + req_idx: Optional[int] = None, + positions: Optional[torch.Tensor] = None, + ): + if kv.shape[0] == 0: + return + packed = self._pack_mla_kv(kv) + if req_idx is None or positions is None: + swa_slots = self._identity_swa_slots(mem_index).to(kv.device) + else: + swa_slots = self.ensure_swa_slots(req_idx, positions, mem_index).to(kv.device) + swa_slots, packed = self._keep_last_swa_writes(swa_slots, packed) + if swa_slots.numel() == 0: + return + self.swa_pool.write(layer_index, swa_slots, packed) + + def pack_compressed_kv_to_cache(self, layer_index: int, slots: torch.Tensor, comp: torch.Tensor): + if comp.shape[0] == 0: + return + pool, local_layer = self._pool_and_local_layer(layer_index) + pool.write_kv(local_layer, slots.to(comp.device), self._pack_mla_kv(comp)) + + def pack_c4_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer_k: torch.Tensor): + if indexer_k.shape[0] == 0: + return + pool, local_layer = self._pool_and_local_layer(layer_index) + pool.write_indexer_k(local_layer, slots.to(indexer_k.device), self._pack_indexer_k(indexer_k)) + + def gather_mla_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + if slots.numel() == 0: + return torch.empty((0, self.mla_head_dim), dtype=self.dtype, device=self.kv_buffer.device) + swa_slots = self._swa_slots_from_full(slots).to(self.kv_buffer.device) + return self._unpack_mla_kv(self.swa_pool.read(layer_index, swa_slots)) + + def gather_compressed_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + if slots.numel() == 0: + return torch.empty((0, self.mla_head_dim), dtype=self.dtype, device=self.kv_buffer.device) + pool, local_layer = self._pool_and_local_layer(layer_index) + return self._unpack_mla_kv(pool.read_kv(local_layer, slots.to(self.kv_buffer.device))) + + def gather_c4_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + if slots.numel() == 0: + return torch.empty( + (0, self.indexer_head_dim), + dtype=self.dtype, + device=self.kv_buffer.device, + ) + pool, local_layer = self._pool_and_local_layer(layer_index) + return self._unpack_indexer_k(pool.read_indexer_k(local_layer, slots.to(self.kv_buffer.device))) def _pool_and_local_layer(self, layer_index: int): r = self.compress_rates[layer_index] @@ -197,7 +704,21 @@ def free_c128(self, free_index) -> None: def free_all(self): super().free_all() + if hasattr(self, "swa_allocator"): + self.swa_allocator.free_all() + if hasattr(self, "full_to_swa_indexs"): + self.full_to_swa_indexs.fill_(-1) + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + if getattr(self, "req_to_swa_indexs", None) is not None: + self.req_to_swa_indexs.fill_(self.swa_pool.HOLD_TOKEN_MEMINDEX) + self.req_to_swa_full_indexs.fill_(-1) if self.c4_pool is not None: self.c4_pool.free_all() if self.c128_pool is not None: self.c128_pool.free_all() + + def alloc_kv_move_buffer(self, max_req_total_len): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") + + def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: + raise NotImplementedError("DeepSeek-V4 packed/composite paged KV transfer is not implemented") diff --git a/lightllm/common/kv_cache_mem_manager/operator/__init__.py b/lightllm/common/kv_cache_mem_manager/operator/__init__.py index 85c37ad39b..442c2e300e 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/__init__.py +++ b/lightllm/common/kv_cache_mem_manager/operator/__init__.py @@ -5,6 +5,7 @@ from .deepseek import ( Deepseek2MemOperator, Deepseek3_2MemOperator, + DeepseekV4MemOperator, FP8PerTokenGroupQuantDeepseek3_2MemOperator, ) from .fp8_quant import ( diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 6e05b96e10..0725ce9b93 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -8,7 +8,9 @@ class Deepseek2MemOperator(NormalMemOperator): def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import Deepseek2MemoryManager + from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import ( + Deepseek2MemoryManager, + ) mem_manager: Deepseek2MemoryManager = self.mem_manager @@ -30,7 +32,9 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: class Deepseek3_2MemOperator(Deepseek2MemOperator): def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - from lightllm.common.kv_cache_mem_manager.deepseek3_2mem_manager import Deepseek3_2MemoryManager + from lightllm.common.kv_cache_mem_manager.deepseek3_2mem_manager import ( + Deepseek3_2MemoryManager, + ) mem_manager: Deepseek3_2MemoryManager = self.mem_manager from ...basemodel.triton_kernel.kv_copy.mla_copy_kv import destindex_copy_kv @@ -78,3 +82,14 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: o_rope, ) return + + +class DeepseekV4MemOperator(BaseMemManagerOperator): + def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DeepseekV4MemoryManager, + ) + + mem_manager: DeepseekV4MemoryManager = self.mem_manager + mem_manager.pack_mla_kv_to_cache(layer_index, mem_index, kv) + return diff --git a/lightllm/common/quantization/__init__.py b/lightllm/common/quantization/__init__.py index cd534d53ec..1c5a9c09d3 100644 --- a/lightllm/common/quantization/__init__.py +++ b/lightllm/common/quantization/__init__.py @@ -68,6 +68,9 @@ def _mapping_quant_method(self): expert_dtype = self.expert_dtype or self.network_config_.get("expert_dtype", None) if expert_dtype is None: return + if expert_dtype == "fp4" and self.network_config_.get("model_type") == "deepseek_v4" and not is_sm100_gpu(): + logger.info("skip generic fused_moe quant mapping for DeepSeek-V4 fp4 experts on non-SM100 GPUs") + return target = self._get_expert_quant_type(expert_dtype) for layer_num in range(self.layer_num): if self.expert_dtype is not None: diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index c8197401c1..1cdea03381 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -6,12 +6,16 @@ from .kv_cache_mem_manager import MemoryManager, DeepseekV4MemoryManager from typing import List, Optional, TYPE_CHECKING from lightllm.common.basemodel.triton_kernel.gen_sampling_params import token_id_counter -from lightllm.common.basemodel.triton_kernel.gen_sampling_params import update_req_to_token_id_counter +from lightllm.common.basemodel.triton_kernel.gen_sampling_params import ( + update_req_to_token_id_counter, +) from lightllm.utils.envs_utils import enable_env_vars, get_env_start_args from lightllm.utils.config_utils import get_vocab_size from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager from lightllm.common.linear_att_cache_manager.layer_cache import LayerCache -from lightllm.common.linear_att_cache_manager.linear_att_buffer_manager import LinearAttCacheManager +from lightllm.common.linear_att_cache_manager.linear_att_buffer_manager import ( + LinearAttCacheManager, +) if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq @@ -131,11 +135,13 @@ def __init__(self, max_request_num): ) elif self.penalty_counter_mode == "pin_mem_counter": self.req_to_out_token_id_counter = torch.zeros( - (max_request_num + 1, self.vocab_size), dtype=torch.int32, device="cpu", pin_memory=True + (max_request_num + 1, self.vocab_size), + dtype=torch.int32, + device="cpu", + pin_memory=True, ) 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()) self.req_to_presence_penalty[req.req_idx].fill_(shm_param.presence_penalty) @@ -165,14 +171,18 @@ def init_req_sampling_params(self, req: "InferReq"): dtype=torch.int32, ).cuda(non_blocking=True) token_id_counter( - prompt_ids=prompt_ids, out_token_id_counter=self.req_to_out_token_id_counter[req.req_idx] + prompt_ids=prompt_ids, + out_token_id_counter=self.req_to_out_token_id_counter[req.req_idx], ) torch.cuda.current_stream().synchronize() return def update_reqs_out_token_counter_gpu( - self, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor, mask: torch.Tensor = None + self, + b_req_idx: torch.Tensor, + next_token_ids: torch.Tensor, + mask: torch.Tensor = None, ): if self.penalty_counter_mode not in ["gpu_counter", "pin_mem_counter"]: return @@ -188,7 +198,10 @@ def update_reqs_out_token_counter_gpu( return def update_reqs_token_counter( - self, req_objs: List["InferReq"], next_token_ids: List[int], accept_mark: Optional[List[List[bool]]] = None + self, + req_objs: List["InferReq"], + next_token_ids: List[int], + accept_mark: Optional[List[List[bool]]] = None, ): if self.penalty_counter_mode != "cpu_counter": return @@ -230,7 +243,13 @@ def gen_cpu_out_token_counter_sampling_params(self, req_objs: List["InferReq"]): class ReqManagerForMamba(ReqManager): - def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_config: LinearAttCacheConfig): + def __init__( + self, + max_request_num, + max_sequence_length, + mem_manager, + linear_config: LinearAttCacheConfig, + ): super().__init__(max_request_num, max_sequence_length, mem_manager) self.mtp_step = get_env_start_args().mtp_step self.big_page_token_num = ( @@ -275,7 +294,6 @@ def get_mamba_cache(self, layer_idx_in_all: int): return conv_states, ssm_states def copy_big_page_buffer_to_linear_att_state(self, big_page_buffer_idx: int, req: "InferReq"): - from .linear_att_cache_manager import LinearAttCacheManager big_page_buffers: LinearAttCacheManager = self.mem_manager.linear_att_big_page_buffers @@ -304,8 +322,9 @@ def copy_small_page_buffer_to_linear_att_state( class DeepseekV4ReqManager(ReqManager): """DeepSeek-V4 的请求级管理(锁定决策: SWA 全历史 + 不分页)。 - 在基类 ReqManager 之上补三类 V4 专有的 per-request 结构(均从 mem_manager 读取 n_c4/n_c128/ - layer_to_*_idx/head_dim 等,避免重复配置): + 在基类 ReqManager 之上补三类 V4 专有的 per-request 结构。该对象在 mem manager profile 前创建, + 所以初始化只依赖 config 派生出的 compress_rates/head_dim/indexer_head_dim;真实 mem_manager + 会在 `_init_mem_manager()` 后通过 `bind_mem_manager()` 接入。 * ``req_to_c4_indexs`` / ``req_to_c128_indexs`` —— (req, 窗口下标) -> 压缩池槽位。 窗口下标 = position // compress_rate;窗口关闭时由 layer-infer 写入,attention 读取前 @@ -318,19 +337,48 @@ class DeepseekV4ReqManager(ReqManager): * entry_count 不另存:= position // compress_rate,可由序列长度推出。 """ - def __init__(self, max_request_num, max_sequence_length, mem_manager: DeepseekV4MemoryManager): + def __init__( + self, + max_request_num, + max_sequence_length, + mem_manager: Optional[DeepseekV4MemoryManager] = None, + compress_rates: Optional[List[int]] = None, + head_dim: Optional[int] = None, + indexer_head_dim: Optional[int] = None, + ): super().__init__(max_request_num, max_sequence_length, mem_manager) - assert isinstance(mem_manager, DeepseekV4MemoryManager) - self.n_c4 = mem_manager.n_c4 - self.n_c128 = mem_manager.n_c128 - head_dim = mem_manager.head_dim - indexer_head_dim = mem_manager.indexer_head_dim + if mem_manager is not None: + assert isinstance(mem_manager, DeepseekV4MemoryManager) + compress_rates = mem_manager.compress_rates + head_dim = mem_manager.head_dim + indexer_head_dim = mem_manager.indexer_head_dim + assert compress_rates is not None, "DeepSeek-V4 req manager requires compress_rates" + assert head_dim is not None, "DeepSeek-V4 req manager requires head_dim" + assert indexer_head_dim is not None, "DeepSeek-V4 req manager requires indexer_head_dim" + + self.compress_rates = list(compress_rates) + self.n_c4 = sum(1 for r in self.compress_rates if r == 4) + self.n_c128 = sum(1 for r in self.compress_rates if r == 128) + self.head_dim = head_dim + self.indexer_head_dim = indexer_head_dim + self.layer_to_c4_idx = {} + self.layer_to_c128_idx = {} + c4 = c128 = 0 + for lid, r in enumerate(self.compress_rates): + if r == 4: + self.layer_to_c4_idx[lid] = c4 + c4 += 1 + elif r == 128: + self.layer_to_c128_idx[lid] = c128 + c128 += 1 # (req, 窗口) -> 压缩槽。列数取 ceil(max_seq / ratio) 留足余量。 c4_windows = (max_sequence_length + 4 - 1) // 4 c128_windows = (max_sequence_length + 128 - 1) // 128 self.req_to_c4_indexs = torch.zeros((max_request_num + 1, c4_windows), dtype=torch.int32, device="cuda") self.req_to_c128_indexs = torch.zeros((max_request_num + 1, c128_windows), dtype=torch.int32, device="cuda") + self._c4_entry_counts = [0 for _ in range(max_request_num + 1)] + self._c128_entry_counts = [0 for _ in range(max_request_num + 1)] # compressor 在途窗口累加状态(fp32): [kv_or_score, coff * ratio, coff * dim]. state_dtype = torch.float32 @@ -355,9 +403,39 @@ def __init__(self, max_request_num, max_sequence_length, mem_manager: DeepseekV4 layer_num=self.n_c4, device="cuda", ) + self.req_to_c4_state_pool = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(1, 8, 4 * head_dim), + layer_num=self.n_c4, + device="cuda", + ) + self.req_to_c128_state_pool = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(1, 128, 2 * head_dim), + layer_num=self.n_c128, + device="cuda", + ) + self.req_to_c4_indexer_state_pool = LayerCache( + size=max_request_num + 1, + dtype=state_dtype, + shape=(1, 8, 4 * indexer_head_dim), + layer_num=self.n_c4, + device="cuda", + ) + self._runtime_states = [{} for _ in range(max_request_num + 1)] self._init_all_score_state() return + def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): + assert isinstance(mem_manager, DeepseekV4MemoryManager) + assert self.compress_rates == mem_manager.compress_rates + assert self.head_dim == mem_manager.head_dim + assert self.indexer_head_dim == mem_manager.indexer_head_dim + self.mem_manager = mem_manager + return + def _init_all_score_state(self): if self.n_c4 > 0: self.req_to_c4_state.buffer[:, :, 1, ...].fill_(float("-inf")) @@ -373,33 +451,184 @@ def _reset_compress_cache_req(self, cache: LayerCache, req_idx: int): cache.buffer[:, req_idx, 1, ...].fill_(float("-inf")) return + def _reset_state_pool_req(self, cache: LayerCache, req_idx: int): + if cache.layer_num == 0: + return + cache.buffer[:, req_idx, ...].fill_(0) + return + def init_compress_state(self, req_idx: int): """新请求开始时重置其 compressor 在途状态(对应 mamba 的 init_linear_att_state)。""" + self.clear_runtime_state(req_idx) + c4, c128 = self.pop_compress_indices_for_req(req_idx) + self.free_compress_indices(free_c4_index=c4, free_c128_index=c128) if self.n_c4 > 0: self._reset_compress_cache_req(self.req_to_c4_state, req_idx) self._reset_compress_cache_req(self.req_to_c4_indexer_state, req_idx) + self._reset_state_pool_req(self.req_to_c4_state_pool, req_idx) + self._reset_state_pool_req(self.req_to_c4_indexer_state_pool, req_idx) if self.n_c128 > 0: self._reset_compress_cache_req(self.req_to_c128_state, req_idx) + self._reset_state_pool_req(self.req_to_c128_state_pool, req_idx) return + def _ensure_compress_slots(self, req_idx: int, ratio: int, entry_start: int, entry_count: int) -> torch.Tensor: + if entry_count == 0: + return torch.empty((0,), dtype=torch.int32, device="cuda") + assert entry_start >= 0 and entry_count >= 0 + assert self.mem_manager is not None, "DeepSeek-V4 mem manager is not bound yet" + if ratio == 4: + table = self.req_to_c4_indexs + counts = self._c4_entry_counts + alloc = self.mem_manager.alloc_c4 + elif ratio == 128: + table = self.req_to_c128_indexs + counts = self._c128_entry_counts + alloc = self.mem_manager.alloc_c128 + else: + raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") + + required_count = entry_start + entry_count + assert required_count <= table.shape[1], ( + f"DeepSeek-V4 compressed slot table overflow: req={req_idx} " + f"ratio={ratio} required={required_count} capacity={table.shape[1]}" + ) + old_count = counts[req_idx] + if required_count > old_count: + new_slots_cpu = alloc(required_count - old_count) + table[req_idx, old_count:required_count] = new_slots_cpu.cuda(non_blocking=True) + counts[req_idx] = required_count + return table[req_idx, entry_start:required_count] + + def ensure_c4_slots(self, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: + return self._ensure_compress_slots(req_idx, 4, entry_start, entry_count) + + def ensure_c128_slots(self, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: + return self._ensure_compress_slots(req_idx, 128, entry_start, entry_count) + + def ensure_compress_slots(self, layer_index: int, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: + ratio = self.compress_rates[layer_index] + if ratio == 4: + return self.ensure_c4_slots(req_idx, entry_start, entry_count) + if ratio == 128: + return self.ensure_c128_slots(req_idx, entry_start, entry_count) + raise AssertionError(f"layer {layer_index} is not a compressed attention layer") + + def pop_compress_indices_for_req(self, req_idx: int): + c4_count = self._c4_entry_counts[req_idx] + if c4_count > 0: + c4 = self.req_to_c4_indexs[req_idx, :c4_count].clone() + self.req_to_c4_indexs[req_idx, :c4_count].fill_(0) + self._c4_entry_counts[req_idx] = 0 + else: + c4 = None + + c128_count = self._c128_entry_counts[req_idx] + if c128_count > 0: + c128 = self.req_to_c128_indexs[req_idx, :c128_count].clone() + self.req_to_c128_indexs[req_idx, :c128_count].fill_(0) + self._c128_entry_counts[req_idx] = 0 + else: + c128 = None + return c4, c128 + + def free_compress_indices(self, free_c4_index=None, free_c128_index=None): + if free_c4_index is not None and len(free_c4_index) > 0: + self.mem_manager.free_c4(free_c4_index) + if free_c128_index is not None and len(free_c128_index) > 0: + self.mem_manager.free_c128(free_c128_index) + return + + def alloc(self): + req_idx = super().alloc() + if req_idx is not None: + self.init_compress_state(req_idx) + return req_idx + + def clear_runtime_state(self, req_idx: int): + self._runtime_states[req_idx].clear() + if self.mem_manager is not None and hasattr(self.mem_manager, "free_swa_for_req"): + self.mem_manager.free_swa_for_req(req_idx) + return + + def set_runtime_state(self, req_idx: int, layer_index: int, state: dict): + self._runtime_states[req_idx][layer_index] = state + return + + def get_runtime_state(self, req_idx: int, layer_index: int): + return self._runtime_states[req_idx][layer_index] + + def get_compress_state_for_req(self, layer_index: int, req_idx: int): + if self.compress_rates[layer_index] == 4: + state = self.get_c4_compress_state(layer_index) + elif self.compress_rates[layer_index] == 128: + state = self.get_c128_compress_state(layer_index) + else: + raise AssertionError(f"layer {layer_index} is not a compressed attention layer") + return state[req_idx, 0], state[req_idx, 1] + + def get_compress_state_pool_for_req(self, layer_index: int, req_idx: int): + if self.compress_rates[layer_index] == 4: + cache = self.req_to_c4_state_pool + local = self.layer_to_c4_idx[layer_index] + elif self.compress_rates[layer_index] == 128: + cache = self.req_to_c128_state_pool + local = self.layer_to_c128_idx[layer_index] + else: + raise AssertionError(f"layer {layer_index} is not a compressed attention layer") + return cache.buffer[local, req_idx] + def get_c4_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.mem_manager.layer_to_c4_idx[layer_index] + local = self.layer_to_c4_idx[layer_index] return self.req_to_c4_state.buffer[local] def get_c128_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.mem_manager.layer_to_c128_idx[layer_index] + local = self.layer_to_c128_idx[layer_index] return self.req_to_c128_state.buffer[local] def get_c4_indexer_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.mem_manager.layer_to_c4_idx[layer_index] + local = self.layer_to_c4_idx[layer_index] return self.req_to_c4_indexer_state.buffer[local] - def free(self, free_req_indexes, free_token_index, free_c4_index=None, free_c128_index=None): + def get_c4_indexer_state_pool_for_req(self, layer_index: int, req_idx: int) -> torch.Tensor: + local = self.layer_to_c4_idx[layer_index] + return self.req_to_c4_indexer_state_pool.buffer[local, req_idx] + + def free( + self, + free_req_indexes, + free_token_index, + free_c4_index=None, + free_c128_index=None, + ): """释放 dense 槽(基类)+ 压缩槽。压缩槽由调用方(infer batch)从 req_to_c*_indexs 收集后传入, 与基类用 free_token_index 传 dense 槽的方式一致。""" + for req_index in free_req_indexes: + self.clear_runtime_state(req_index) super().free(free_req_indexes, free_token_index) - if free_c4_index is not None and len(free_c4_index) > 0: - self.mem_manager.free_c4(free_c4_index) - if free_c128_index is not None and len(free_c128_index) > 0: - self.mem_manager.free_c128(free_c128_index) + self.free_compress_indices(free_c4_index=free_c4_index, free_c128_index=free_c128_index) + return + + def free_req(self, free_req_index: int): + self.clear_runtime_state(free_req_index) + c4, c128 = self.pop_compress_indices_for_req(free_req_index) + self.free_compress_indices(free_c4_index=c4, free_c128_index=c128) + return super().free_req(free_req_index) + + def free_all(self): + super().free_all() + self._runtime_states = [{} for _ in range(self.max_request_num + 1)] + self._c4_entry_counts = [0 for _ in range(self.max_request_num + 1)] + self._c128_entry_counts = [0 for _ in range(self.max_request_num + 1)] + if self.n_c4 > 0: + self.req_to_c4_indexs.fill_(0) + self.req_to_c4_state.buffer.fill_(0) + self.req_to_c4_indexer_state.buffer.fill_(0) + self.req_to_c4_state_pool.buffer.fill_(0) + self.req_to_c4_indexer_state_pool.buffer.fill_(0) + if self.n_c128 > 0: + self.req_to_c128_indexs.fill_(0) + self.req_to_c128_state.buffer.fill_(0) + self.req_to_c128_state_pool.buffer.fill_(0) + self._init_all_score_state() return diff --git a/lightllm/models/deepseek3_2/model.py b/lightllm/models/deepseek3_2/model.py index 5831044311..cd33386666 100644 --- a/lightllm/models/deepseek3_2/model.py +++ b/lightllm/models/deepseek3_2/model.py @@ -1,14 +1,20 @@ import copy from lightllm.models.registry import ModelRegistry from lightllm.models.deepseek2.model import Deepseek2TpPartModel -from lightllm.models.deepseek3_2.layer_weights.transformer_layer_weight import Deepseek3_2TransformerLayerWeight -from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer -from lightllm.common.basemodel.attention import get_nsa_prefill_att_backend_class, get_nsa_decode_att_backend_class +from lightllm.models.deepseek3_2.layer_weights.transformer_layer_weight import ( + Deepseek3_2TransformerLayerWeight, +) +from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import ( + Deepseek3_2TransformerLayerInfer, +) +from lightllm.common.basemodel.attention import ( + get_nsa_prefill_att_backend_class, + get_nsa_decode_att_backend_class, +) @ModelRegistry(["deepseek_v32"]) class Deepseek3_2TpPartModel(Deepseek2TpPartModel): - # weight class transformer_weight_class = Deepseek3_2TransformerLayerWeight @@ -21,24 +27,11 @@ def _init_att_backend(self): return -class DeepSeekV32Tokenizer: - """Tokenizer wrapper for DeepSeek-V3.2 that uses the Python-based - encoding_dsv32 module instead of Jinja chat templates. - - DeepSeek-V3.2's tokenizer_config.json does not ship with a Jinja chat - template, so ``apply_chat_template`` would fail without either a manually - supplied ``--chat_template`` file or this wrapper. - """ - +class DeepSeekChatTokenizerBase: def __init__(self, tokenizer): self.tokenizer = tokenizer - # Cache added vocabulary for performance (HuggingFace can be slow). self._added_vocab = None - # ------------------------------------------------------------------ - # Attribute delegation – everything not overridden goes to the inner - # tokenizer so that encode/decode/vocab_size/eos_token_id/… all work. - # ------------------------------------------------------------------ def __getattr__(self, name): return getattr(self.tokenizer, name) @@ -47,9 +40,9 @@ def get_added_vocab(self): self._added_vocab = self.tokenizer.get_added_vocab() return self._added_vocab - # ------------------------------------------------------------------ - # Core override: route apply_chat_template through encode_messages. - # ------------------------------------------------------------------ + def _encode_messages(self, msgs, thinking_mode, kwargs): + raise NotImplementedError("subclass must provide DeepSeek encode_messages") + def apply_chat_template( self, conversation=None, @@ -58,27 +51,16 @@ def apply_chat_template( tokenize=False, add_generation_prompt=True, thinking=None, + enable_thinking=None, **kwargs, ): - from lightllm.models.deepseek3_2.encoding_dsv32 import encode_messages, render_tools - msgs = conversation if conversation is not None else messages if msgs is None: raise ValueError("Either 'conversation' or 'messages' must be provided") - # Deep copy to avoid mutating the caller's messages. msgs = copy.deepcopy(msgs) - # Determine thinking mode. - thinking_mode = "thinking" if thinking else "chat" - - # Inject tools into the first system message (or create one) so that - # encode_messages / render_message picks them up. if tools: - # build_prompt passes tools as bare function dicts: - # [{"name": "f", "description": "...", "parameters": {...}}] - # encoding_dsv32's render_message expects OpenAI wrapper format: - # [{"type": "function", "function": {...}}] wrapped_tools = [] for t in tools: if "function" in t: @@ -95,16 +77,27 @@ def apply_chat_template( break if not injected: - # Prepend a system message that carries the tools. msgs.insert(0, {"role": "system", "content": "", "tools": wrapped_tools}) - prompt = encode_messages( + if thinking is None: + thinking = bool(enable_thinking) if enable_thinking is not None else False + thinking_mode = "thinking" if thinking else "chat" + prompt = self._encode_messages(msgs, thinking_mode, kwargs) + + if tokenize: + return self.tokenizer.encode(prompt, add_special_tokens=False) + return prompt + + +class DeepSeekV32Tokenizer(DeepSeekChatTokenizerBase): + """Tokenizer wrapper for DeepSeek-V3.2's Python-based encoding_dsv32 module.""" + + def _encode_messages(self, msgs, thinking_mode, kwargs): + from lightllm.models.deepseek3_2.encoding_dsv32 import encode_messages + + return encode_messages( msgs, thinking_mode=thinking_mode, drop_thinking=kwargs.get("drop_thinking", True), add_default_bos_token=kwargs.get("add_default_bos_token", True), ) - - if tokenize: - return self.tokenizer.encode(prompt, add_special_tokens=False) - return prompt diff --git a/lightllm/models/deepseek_v4/layer_infer/attention.py b/lightllm/models/deepseek_v4/layer_infer/attention.py index a25a2aa3d1..a24949696f 100644 --- a/lightllm/models/deepseek_v4/layer_infer/attention.py +++ b/lightllm/models/deepseek_v4/layer_infer/attention.py @@ -1,34 +1,101 @@ +import os + import torch -import torch.nn.functional as F -# DeepSeek-V4 attention: MLA with a single shared KV head (head_dim=512), per-head learnable attention -# sink, and a candidate set = sliding-window tokens (size `window`) ++ compressed KV entries. Pure-torch -# transcription of the bundled reference (inference/model.py Attention.forward + kernel.py sparse_attn). -# Correctness-first prefill path. head_dim=512 > 256 so FlashAttention is unusable anyway; a fused -# triton sparse-gather kernel is a perf follow-up. +FLASHMLA_MIN_HEADS = 64 +FLASHMLA_TOPK_MULTIPLE = 128 +DSV4_DEBUG_TORCH_SPARSE_ATTN = os.getenv("DSV4_DEBUG_TORCH_SPARSE_ATTN", "0") == "1" + + +def _pad_topk_for_flashmla(topk_idxs): + K = topk_idxs.shape[-1] + padded_K = ((K + FLASHMLA_TOPK_MULTIPLE - 1) // FLASHMLA_TOPK_MULTIPLE) * FLASHMLA_TOPK_MULTIPLE + if padded_K == K: + return topk_idxs.contiguous() + padded = torch.full((*topk_idxs.shape[:-1], padded_K), -1, device=topk_idxs.device, dtype=topk_idxs.dtype) + padded[..., :K] = topk_idxs + return padded.contiguous() + + +def _compact_topk_indices(topk_idxs, kv_len): + valid = (topk_idxs >= 0) & (topk_idxs < kv_len) + topk_lens = valid.sum(dim=-1).to(torch.int32) + if valid.all(): + return topk_idxs.contiguous(), topk_lens.contiguous() + + compact = torch.full_like(topk_idxs, -1) + ranks = valid.to(torch.int32).cumsum(dim=-1) - 1 + rows = torch.arange(topk_idxs.shape[0], device=topk_idxs.device).unsqueeze(1).expand_as(topk_idxs) + compact[rows[valid], ranks[valid].long()] = topk_idxs[valid] + return compact.contiguous(), topk_lens.contiguous() + + +def _pad_heads_for_flashmla(q, attn_sink): + h = q.shape[1] + if h == FLASHMLA_MIN_HEADS: + return q.contiguous(), attn_sink.to(torch.float32).contiguous(), h + if h > FLASHMLA_MIN_HEADS: + raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention only supports up to 64 local heads, got {h}") -def torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): - """Gather-then-softmax attention with a per-head sink, matching reference kernel.sparse_attn. + q_pad = q.new_zeros(q.shape[0], FLASHMLA_MIN_HEADS, q.shape[2]) + q_pad[:, :h] = q + sink_pad = torch.full((FLASHMLA_MIN_HEADS,), -float("inf"), device=q.device, dtype=torch.float32) + sink_pad[:h] = attn_sink.to(torch.float32) + return q_pad.contiguous(), sink_pad.contiguous(), h - q:[b,m,h,d], kv:[b,n,d] (single KV head shared over h), attn_sink:[h] (fp32), - topk_idxs:[b,m,K] int (-1 = invalid/skip). Returns o:[b,m,h,d]. + +def _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): + q0 = q[0].float() + kv0 = kv[0].float() + indices = topk_idxs[0].long() + valid = (indices >= 0) & (indices < kv0.shape[0]) + safe_indices = torch.where(valid, indices, torch.zeros_like(indices)) + kv_sel = kv0[safe_indices] + scores = torch.einsum("mhd,mkd->mhk", q0, kv_sel) * scale + scores = scores.masked_fill(~valid.unsqueeze(1), float("-inf")) + sink = attn_sink.float().view(1, -1) + max_scores = torch.maximum(scores.max(dim=-1).values, sink) + exp_scores = torch.exp(scores - max_scores.unsqueeze(-1)).masked_fill(~valid.unsqueeze(1), 0.0) + exp_sink = torch.exp(sink - max_scores) + denom = exp_scores.sum(dim=-1) + exp_sink + out = torch.einsum("mhk,mkd->mhd", exp_scores / denom.unsqueeze(-1), kv_sel) + return out.unsqueeze(0).to(q.dtype) + + +def vllm_sparse_attn(q, kv, attn_sink, topk_idxs, scale): + """DeepSeek-V4 sparse MLA through vLLM FlashMLA. + + q:[1,m,h,d], kv:[1,n,d] (single KV head shared over h), attn_sink:[h], + topk_idxs:[1,m,K] int (-1 = invalid/skip). Returns o:[1,m,h,d]. """ b, m, h, d = q.shape - n = kv.shape[1] - K = topk_idxs.shape[-1] - idx = topk_idxs.clamp(min=0).long() # [b,m,K] - keys = torch.gather(kv.unsqueeze(1).expand(b, m, n, d), 2, idx.unsqueeze(-1).expand(b, m, K, d)) # [b,m,K,d] - qf, kf = q.float(), keys.float() - scores = torch.einsum("bmhd,bmkd->bmhk", qf, kf) * scale # [b,m,h,K] - valid = (topk_idxs != -1).unsqueeze(2) # [b,m,1,K] - scores = scores.masked_fill(~valid, float("-inf")) - mx = scores.amax(dim=-1, keepdim=True) # [b,m,h,1] - mx = torch.nan_to_num(mx, neginf=0.0) - ex = (scores - mx).exp() # [b,m,h,K] - denom = ex.sum(-1) + (attn_sink.view(1, 1, h) - mx.squeeze(-1)).exp() # [b,m,h] - o = torch.einsum("bmhk,bmkd->bmhd", ex, kf) / denom.unsqueeze(-1) - return o.to(q.dtype) + if b != 1 or kv.shape[0] != 1 or topk_idxs.shape[0] != 1: + raise RuntimeError("DeepSeek-V4 FlashMLA sparse attention wrapper expects one request per call") + if d != 512: + raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires head_dim=512, got {d}") + if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: + raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires bf16 q/kv, got {q.dtype}/{kv.dtype}") + + if DSV4_DEBUG_TORCH_SPARSE_ATTN: + return _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale) + + from vllm.third_party.flashmla.flash_mla_interface import flash_mla_sparse_fwd + + q_pad, sink_pad, real_heads = _pad_heads_for_flashmla(q[0], attn_sink) + indices, topk_lens = _compact_topk_indices(topk_idxs[0].to(torch.int32), kv.shape[1]) + indices = _pad_topk_for_flashmla(indices).unsqueeze(1) + kv_flat = kv[0].unsqueeze(1).contiguous() + out, _, _ = flash_mla_sparse_fwd( + q=q_pad, + kv=kv_flat, + indices=indices, + sm_scale=scale, + attn_sink=sink_pad, + topk_length=topk_lens, + out=None, + ) + return out[:, :real_heads].unsqueeze(0).to(q.dtype) def build_prefill_topk_idxs(seqlen, window, ratio, n_window, device): @@ -40,11 +107,9 @@ def build_prefill_topk_idxs(seqlen, window, ratio, n_window, device): entries are attended (matches the reference for short context). """ t = torch.arange(seqlen, device=device) - # sliding window: query t attends tokens [max(0, t-window+1) .. t] - j = torch.arange(n_window, device=device) - win = j.unsqueeze(0).expand(seqlen, n_window).clone() # [s, n_window] - win_valid = (j.unsqueeze(0) <= t.unsqueeze(1)) & (j.unsqueeze(0) > (t.unsqueeze(1) - window)) - win = torch.where(win_valid, win, torch.full_like(win, -1)) + offsets = torch.arange(window, device=device) + win = t.unsqueeze(1) - (window - 1 - offsets).unsqueeze(0) + win = torch.where(win >= 0, win, torch.full_like(win, -1)) if ratio: ncomp = seqlen // ratio c = torch.arange(ncomp, device=device) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 902de113db..c91799f9ee 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -1,7 +1,20 @@ +import importlib.util +import logging +import sys +import types +from pathlib import Path + import torch import torch.nn.functional as F from ..triton_kernel.rotary_emb import apply_rotary_emb +logger = logging.getLogger(__name__) + +_SGLANG_COMPRESS_MOD = None +_SGLANG_COMPRESS_ERR = None +_SGLANG_COMPRESS_WARNED = False +_FREQ_CIS_CACHE = {} + # KV compressor: pools every `ratio` consecutive tokens into one compressed KV entry via gated # (softmax) pooling + a learned absolute-position bias (ape), RMSNorm, and rope on the trailing # rope_dim. ratio==4 uses overlapping windows (two-series Ca/Cb scheme). Pure-torch transcription of @@ -25,6 +38,231 @@ def _rmsnorm(x, weight, eps): return (xf * weight.float()).to(x.dtype) +def _load_file_module(name, path): + spec = importlib.util.spec_from_file_location(name, path) + mod = importlib.util.module_from_spec(spec) + sys.modules[name] = mod + spec.loader.exec_module(mod) + return mod + + +def _load_sglang_compressor(): + global _SGLANG_COMPRESS_MOD, _SGLANG_COMPRESS_ERR + if _SGLANG_COMPRESS_MOD is not None: + return _SGLANG_COMPRESS_MOD + if _SGLANG_COMPRESS_ERR is not None: + raise _SGLANG_COMPRESS_ERR + try: + from sglang.jit_kernel.dsv4 import compress_old as mod + + _SGLANG_COMPRESS_MOD = mod + return mod + except Exception as first_exc: + root = Path("/data/wanzihao/sglang/python/sglang") + try: + if not root.exists(): + raise first_exc + if "sglang" not in sys.modules: + sglang_mod = types.ModuleType("sglang") + sglang_mod.__path__ = [str(root)] + sys.modules["sglang"] = sglang_mod + if "sglang.utils" not in sys.modules: + utils_mod = types.ModuleType("sglang.utils") + utils_mod.is_in_ci = lambda: False + sys.modules["sglang.utils"] = utils_mod + if "sglang.jit_kernel" not in sys.modules: + jit_mod = types.ModuleType("sglang.jit_kernel") + jit_mod.__path__ = [str(root / "jit_kernel")] + sys.modules["sglang.jit_kernel"] = jit_mod + if "sglang.jit_kernel.dsv4" not in sys.modules: + dsv4_mod = types.ModuleType("sglang.jit_kernel.dsv4") + dsv4_mod.__path__ = [str(root / "jit_kernel" / "dsv4")] + sys.modules["sglang.jit_kernel.dsv4"] = dsv4_mod + if "sglang.srt" not in sys.modules: + srt_mod = types.ModuleType("sglang.srt") + srt_mod.__path__ = [str(root / "srt")] + sys.modules["sglang.srt"] = srt_mod + if "sglang.srt.environ" not in sys.modules: + env_mod = types.ModuleType("sglang.srt.environ") + + class _FalseEnv: + def get(self): + return False + + class _Envs: + SGLANG_OPT_USE_ONLINE_COMPRESS = _FalseEnv() + + env_mod.envs = _Envs() + sys.modules["sglang.srt.environ"] = env_mod + if "sglang.jit_kernel.utils" not in sys.modules: + _load_file_module("sglang.jit_kernel.utils", root / "jit_kernel" / "utils.py") + if "sglang.jit_kernel.dsv4.utils" not in sys.modules: + _load_file_module( + "sglang.jit_kernel.dsv4.utils", + root / "jit_kernel" / "dsv4" / "utils.py", + ) + _SGLANG_COMPRESS_MOD = _load_file_module( + "sglang.jit_kernel.dsv4.compress_old", + root / "jit_kernel" / "dsv4" / "compress_old.py", + ) + return _SGLANG_COMPRESS_MOD + except Exception as exc: + _SGLANG_COMPRESS_ERR = exc + raise exc + + +def _warn_sglang_fallback(exc): + global _SGLANG_COMPRESS_WARNED + if not _SGLANG_COMPRESS_WARNED: + logger.warning("DeepSeek-V4 SGLang compressor JIT unavailable, fallback to torch: %s", exc) + _SGLANG_COMPRESS_WARNED = True + + +def _freq_cis(cos_table, sin_table): + key = ( + cos_table.data_ptr(), + sin_table.data_ptr(), + cos_table.device, + tuple(cos_table.shape), + tuple(sin_table.shape), + ) + cached = _FREQ_CIS_CACHE.get(key) + if cached is None: + cached = torch.complex(cos_table.float(), sin_table.float()) + _FREQ_CIS_CACHE[key] = cached + return cached + + +def _sglang_ape(ape, ratio, head_dim): + if ratio == 4: + return torch.cat([ape[:, :head_dim], ape[:, head_dim:]], dim=0).contiguous() + return ape.contiguous() + + +def _pack_kv_score(kv, score, ratio, head_dim): + if ratio == 4: + return torch.cat( + [ + kv[:, :head_dim], + kv[:, head_dim:], + score[:, :head_dim], + score[:, head_dim:], + ], + dim=1, + ).contiguous() + return torch.cat([kv, score], dim=1).contiguous() + + +def _build_state_from_kv_score(kv, score, ape, ratio, head_dim): + overlap = ratio == 4 + kv_state, score_state = new_compressor_state(ratio, head_dim, kv.device) + s = kv.shape[0] + remainder = s % ratio + cutoff = s - remainder + offset = ratio if overlap else 0 + if overlap and cutoff >= ratio: + kv_state[:ratio] = kv[cutoff - ratio : cutoff] + score_state[:ratio] = score[cutoff - ratio : cutoff] + ape.float() + if remainder > 0: + kv_state[offset : offset + remainder] = kv[cutoff:] + score_state[offset : offset + remainder] = score[cutoff:] + ape.float()[:remainder] + return kv_state, score_state + + +def _sglang_prefill_from_kv_score( + kv, + score, + norm_w, + ape, + ratio, + head_dim, + cos_table, + sin_table, + eps, + dtype, + state_pool=None, +): + if not kv.is_cuda or head_dim % 128 != 0 or ratio not in (4, 128): + return None, None + mod = _load_sglang_compressor() + kv_score = _pack_kv_score(kv, score, ratio, head_dim) + ape_sglang = _sglang_ape(ape.float(), ratio, head_dim) + slots = 8 if ratio == 4 else ratio + if state_pool is None: + state_pool = torch.zeros((1, slots, kv_score.shape[1]), device=kv.device, dtype=kv_score.dtype) + else: + state_pool.zero_() + seq_len = kv.shape[0] + plan = mod.CompressorPrefillPlan.generate( + ratio, + seq_len, + torch.tensor([seq_len], dtype=torch.int64), + torch.tensor([seq_len], dtype=torch.int64), + kv.device, + ) + indices = torch.zeros((1,), device=kv.device, dtype=torch.int32) + out = mod.compress_forward( + state_pool, + kv_score, + ape_sglang, + indices, + plan, + head_dim=head_dim, + compress_ratio=ratio, + ) + ncomp = seq_len // ratio + if ncomp: + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() + out = out.index_select(0, ragged_ids).to(dtype) + else: + out = kv.new_zeros(0, head_dim).to(dtype) + return out, state_pool + + +def _sglang_decode_step_from_state_pool( + x_new, + wkv_w, + wgate_w, + norm_w, + ape, + ratio, + head_dim, + cos_table, + sin_table, + eps, + start_pos, + state_pool, +): + if state_pool is None or not x_new.is_cuda or head_dim % 128 != 0 or ratio not in (4, 128): + return None, False + mod = _load_sglang_compressor() + xf = x_new.float().view(1, -1) + kv = F.linear(xf, wkv_w.float()) + score = F.linear(xf, wgate_w.float()) + kv_score = _pack_kv_score(kv, score, ratio, head_dim) + ape_sglang = _sglang_ape(ape.float(), ratio, head_dim) + seq_len = start_pos + 1 + plan = mod.CompressorDecodePlan( + ratio, + torch.tensor([seq_len], device=x_new.device, dtype=torch.int32), + ) + indices = torch.zeros((1,), device=x_new.device, dtype=torch.int32) + out = mod.compress_forward( + state_pool, + kv_score, + ape_sglang, + indices, + plan, + head_dim=head_dim, + compress_ratio=ratio, + ) + if seq_len % ratio != 0: + return None, True + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + return out[0].to(x_new.dtype), True + + def compress_prefill(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, rope_dim, cos_table, sin_table, eps): """x:[s,dim] (one request, start_pos=0) -> compressed kv [nwin, head_dim] (rope applied to last rope_dim). @@ -69,7 +307,21 @@ def _finish_entry(kv, norm_w, ape_unused, rope_dim, cos_table, sin_table, positi return torch.cat([kv[:-rope_dim], kv_rope], dim=0) -def compressor_prefill_state(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, rope_dim, cos_table, sin_table, eps): +def compressor_prefill_state( + x, + wkv_w, + wgate_w, + norm_w, + ape, + ratio, + head_dim, + rope_dim, + cos_table, + sin_table, + eps, + return_state_pool=False, + state_pool=None, +): """Faithful reference start_pos==0 path (incl. remainder). Returns (entries[ncomp,d], kv_state, score_state). entries have rope applied; kv_state/score_state carry the partial window for the decode path. @@ -83,21 +335,40 @@ def compressor_prefill_state(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, ro kv = F.linear(xf, wkv_w.float()) # [s, coff*d] score = F.linear(xf, wgate_w.float()) # [s, coff*d] ape = ape.float() - kv_state, score_state = new_compressor_state(ratio, head_dim, x.device) + kv_state, score_state = _build_state_from_kv_score(kv, score, ape, ratio, head_dim) + sglang_state_pool = state_pool + try: + comp, sglang_state_pool = _sglang_prefill_from_kv_score( + kv, + score, + norm_w, + ape, + ratio, + head_dim, + cos_table, + sin_table, + eps, + dtype, + state_pool=sglang_state_pool, + ) + if comp is not None: + if return_state_pool: + return comp, kv_state, score_state, sglang_state_pool + return comp, kv_state, score_state + except Exception as exc: + _warn_sglang_fallback(exc) + should_compress = s >= ratio remainder = s % ratio cutoff = s - remainder - offset = ratio if overlap else 0 - if overlap and cutoff >= ratio: - kv_state[:ratio] = kv[cutoff - ratio : cutoff] - score_state[:ratio] = score[cutoff - ratio : cutoff] + ape if remainder > 0: - kv_state[offset : offset + remainder] = kv[cutoff:] - score_state[offset : offset + remainder] = score[cutoff:] + ape[:remainder] kv = kv[:cutoff] score = score[:cutoff] if not should_compress: - return x.new_zeros(0, head_dim), kv_state, score_state + comp = x.new_zeros(0, head_dim) + if return_state_pool: + return comp, kv_state, score_state, sglang_state_pool + return comp, kv_state, score_state nwin = cutoff // ratio kvw = kv.view(nwin, ratio, coff * d) scw = score.view(nwin, ratio, coff * d) + ape @@ -109,6 +380,8 @@ def compressor_prefill_state(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, ro pos = torch.arange(nwin, device=x.device) * ratio comp_rope = apply_rotary_emb(comp[:, -rope_dim:], cos_table[pos], sin_table[pos]) comp = torch.cat([comp[:, :-rope_dim], comp_rope], dim=1) + if return_state_pool: + return comp, kv_state, score_state, sglang_state_pool return comp, kv_state, score_state @@ -127,12 +400,34 @@ def compressor_decode_step( kv_state, score_state, start_pos, + state_pool=None, ): """Faithful reference start_pos>0 path for one new token. Mutates kv_state/score_state in place. - Returns the new compressed entry [d] (rope applied) when a window completes, else None.""" + Returns the new compressed entry [d] (rope applied) when a window completes, else None. + """ overlap = ratio == 4 d = head_dim dtype = x_new.dtype + try: + entry, handled = _sglang_decode_step_from_state_pool( + x_new, + wkv_w, + wgate_w, + norm_w, + ape, + ratio, + head_dim, + cos_table, + sin_table, + eps, + start_pos, + state_pool, + ) + if handled: + return entry + except Exception as exc: + _warn_sglang_fallback(exc) + xf = x_new.float().view(-1) # [dim] kv = F.linear(xf, wkv_w.float()) # [coff*d] score = F.linear(xf, wgate_w.float()) + ape.float()[start_pos % ratio] # [coff*d] @@ -153,4 +448,14 @@ def compressor_decode_step( entry = (kv_state * torch.softmax(score_state, dim=0)).sum(dim=0) # [d] if not should_compress: return None - return _finish_entry(entry, norm_w, ape, rope_dim, cos_table, sin_table, start_pos + 1 - ratio, eps, dtype) + return _finish_entry( + entry, + norm_w, + ape, + rope_dim, + cos_table, + sin_table, + start_pos + 1 - ratio, + eps, + dtype, + ) diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index 75f540725b..78cdb3a3f8 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -1,58 +1,50 @@ import torch -import torch.nn.functional as F - -# Manifold-constrained Hyper-Connections (mHC). Replaces the plain residual add: the hidden state is -# carried as ``hc_mult`` parallel streams. Each sub-layer (attn / ffn) collapses the streams to one -# vector (hc_pre), runs the sub-layer, then re-expands into the streams via learned post/comb weights -# (hc_post). A doubly-stochastic (Sinkhorn-normalized) ``comb`` matrix mixes the residual streams. -# Pure-torch transcription of the bundled reference inference/model.py (Block.hc_pre/hc_post, -# ParallelHead.hc_head) + inference/kernel.py (hc_split_sinkhorn). All math in fp32, as in the reference. - - -def hc_split_sinkhorn(mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps): - """mixes:[N, (2+hc)*hc] fp32 -> pre[N,hc], post[N,hc], comb[N,hc,hc] (doubly stochastic).""" - hc = hc_mult - pre = torch.sigmoid(mixes[:, :hc] * hc_scale[0] + hc_base[:hc]) + eps - post = 2.0 * torch.sigmoid(mixes[:, hc : 2 * hc] * hc_scale[1] + hc_base[hc : 2 * hc]) - comb = mixes[:, 2 * hc :].view(-1, hc, hc) * hc_scale[2] + hc_base[2 * hc :].view(hc, hc) - # comb = softmax(comb, dim=-1) + eps - comb = torch.softmax(comb, dim=-1) + eps - # one column normalization, then (iters-1) of (row, column) - comb = comb / (comb.sum(dim=-2, keepdim=True) + eps) - for _ in range(sinkhorn_iters - 1): - comb = comb / (comb.sum(dim=-1, keepdim=True) + eps) - comb = comb / (comb.sum(dim=-2, keepdim=True) + eps) - return pre, post, comb + + +def _ensure_vllm_mhc_ops(): + try: + import vllm.model_executor.layers.mhc # noqa: F401 + except Exception as e: + raise RuntimeError("DeepSeek-V4 requires vLLM mHC custom ops; failed to import vllm MHC kernels") from e def hc_pre(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps, sinkhorn_iters): - """streams:[N, hc*dim] -> (collapsed[N,dim], post[N,hc], comb[N,hc,hc]).""" - dtype = streams.dtype - x = streams.float() # [N, hc*dim] - rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + eps) - mixes = F.linear(x, hc_fn) * rsqrt # [N, (2+hc)*hc] - pre, post, comb = hc_split_sinkhorn(mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps) - streams3 = x.view(-1, hc_mult, dim) - collapsed = torch.sum(pre.unsqueeze(-1) * streams3, dim=1) # [N, dim] - return collapsed.to(dtype), post, comb + """streams:[N, hc*dim] -> (collapsed[N,dim], post[N,hc,1], comb[N,hc,hc]).""" + _ensure_vllm_mhc_ops() + post, comb, collapsed = torch.ops.vllm.mhc_pre( + residual=streams.view(-1, hc_mult, dim).contiguous(), + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=eps, + hc_pre_eps=eps, + hc_sinkhorn_eps=eps, + hc_post_mult_value=2.0, + sinkhorn_repeat=sinkhorn_iters, + ) + return collapsed, post, comb def hc_post(x, residual, post, comb, hc_mult, dim): """x:[N,dim] sub-layer output, residual:[N, hc*dim] -> [N, hc*dim].""" - res = residual.float().view(-1, hc_mult, dim) # [N, hc, dim] - xf = x.float() - # post: [N,hc] -> [N,hc,dim]; comb mixes residual streams: out[i] = post[i]*x + sum_j comb[i,j]*res[j] - y = post.unsqueeze(-1) * xf.unsqueeze(-2) + torch.einsum("nij,njd->nid", comb, res) - return y.reshape(-1, hc_mult * dim).to(x.dtype) + _ensure_vllm_mhc_ops() + out = torch.ops.vllm.mhc_post(x, residual.view(-1, hc_mult, dim).contiguous(), post, comb) + return out.reshape(-1, hc_mult * dim) def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps): - """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim] (sigmoid gate, no sinkhorn).""" - dtype = streams.dtype - x = streams.float() - rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + eps) - mixes = F.linear(x, hc_fn) * rsqrt # [N, hc] - pre = torch.sigmoid(mixes * hc_scale + hc_base) + eps # [N, hc] - streams3 = x.view(-1, hc_mult, dim) - collapsed = torch.sum(pre.unsqueeze(-1) * streams3, dim=1) - return collapsed.to(dtype) + """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim].""" + _ensure_vllm_mhc_ops() + out = torch.empty(streams.shape[0], dim, device=streams.device, dtype=streams.dtype) + torch.ops.vllm.hc_head_fused_kernel( + streams.view(-1, hc_mult, dim).contiguous(), + hc_fn, + hc_scale, + hc_base, + out, + dim, + eps, + eps, + hc_mult, + ) + return out diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index a8dd0bb1e8..98dee7fd8a 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,3 +1,5 @@ +import os + import torch import torch.nn.functional as F import torch.distributed as dist @@ -7,21 +9,23 @@ from .hyper_connection import hc_pre, hc_post from ..triton_kernel.rotary_emb import apply_rotary_emb from .compressor import compressor_prefill_state, compressor_decode_step -from .attention import torch_sparse_attn -from ..triton_kernel.quant_convert import dequant_fp4_group_to_bf16 +from .attention import vllm_sparse_attn + + +DSV4_DEBUG_DIRECT_PREFILL_COMP = os.getenv("DSV4_DEBUG_DIRECT_PREFILL_COMP", "0") == "1" +DSV4_DEBUG_DISABLE_COMP_ATTN = os.getenv("DSV4_DEBUG_DISABLE_COMP_ATTN", "0") == "1" class DeepseekV4TransformerLayerInfer(TransformerLayerInferTpl): - """One V4 decoder layer: HC(attn) then HC(ffn). Correctness-first pure-torch. + """One V4 decoder layer: HC(attn) then HC(ffn). The residual is carried as ``hc_mult`` streams flattened to [T, hc_mult*hidden]; each sub-layer collapses (hc_pre), computes, and re-expands (hc_post). Attention is MLA over a sliding window + - compressed KV with a per-head sink (torch_sparse_attn); the MoE reuses lightllm's deepgemm FP8 + compressed KV with a per-head sink (vLLM FlashMLA sparse); the MoE reuses lightllm's deepgemm FP8 grouped GEMM driven by V4's custom router (sqrtsoftplus + hash/topk + bias-for-selection). Per-request decode state (window KV history + compressed KV + compressor running state) is kept in - a dict keyed by request id. NOTE: correctness-first — this should move into the KV mem manager for - production memory management / request eviction. + DeepseekV4ReqManager so request alloc/free owns its lifetime. """ def __init__(self, layer_num, network_config): @@ -32,6 +36,9 @@ def __init__(self, layer_num, network_config): self.n_heads = cfg["num_attention_heads"] self.head_dim = cfg["head_dim"] self.rope_dim = cfg["qk_rope_head_dim"] + self.index_n_heads = cfg["index_n_heads"] + self.index_head_dim = cfg["index_head_dim"] + self.index_topk = cfg["index_topk"] self.o_groups = cfg["o_groups"] self.o_lora = cfg["o_lora_rank"] self.hc_mult = cfg["hc_mult"] @@ -43,12 +50,14 @@ def __init__(self, layer_num, network_config): self.topk = cfg["num_experts_per_tok"] self.route_scale = cfg["routed_scaling_factor"] self.swiglu_limit = cfg["swiglu_limit"] - self.softmax_scale = self.head_dim**-0.5 + self.softmax_scale = self.head_dim ** -0.5 self.tp_q_heads = self.n_heads // self.tp_world_size_ + self.tp_index_heads = self.index_n_heads // self.tp_world_size_ self.tp_groups = self.o_groups // self.tp_world_size_ self.embed_dim_ = self.hc_mult * self.hidden self.enable_ep_moe = get_env_start_args().enable_ep_moe - self._state = {} # req_id -> dict(kv_hist, comp_kv, cstate_kv, cstate_score) + self.indexer_score_scale = self.index_head_dim ** -0.5 + self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 # ------------------------------------------------------------------ forward (HC-wrapped) def _hc_block(self, streams, infer_state, lw, attn_fn): @@ -100,8 +109,14 @@ def _qkv(self, x, cos_tok, sin_tok, lw): dim=-1, ) kv = lw.kv_norm_(lw.wkv_.mm(x), eps=self.eps_) - kv = torch.cat([kv[:, : -self.rope_dim], apply_rotary_emb(kv[:, -self.rope_dim :], cos_tok, sin_tok)], dim=1) - return q, kv + kv = torch.cat( + [ + kv[:, : -self.rope_dim], + apply_rotary_emb(kv[:, -self.rope_dim :], cos_tok, sin_tok), + ], + dim=1, + ) + return q, kv, qa def _out_proj(self, o, infer_state, lw): # o: [T, tp_q_heads, head_dim] -> inverse rope -> grouped low-rank O -> [T, hidden] @@ -117,35 +132,130 @@ def _inv_rope(self, o, cos_tok, sin_tok): return torch.cat( [ o[..., : -self.rope_dim], - apply_rotary_emb(o[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1), inverse=True), + apply_rotary_emb( + o[..., -self.rope_dim :], + cos_tok.unsqueeze(1), + sin_tok.unsqueeze(1), + inverse=True, + ), ], dim=-1, ) + def _post_dense_kv(self, infer_state, req, start_pos, mem_index, kv): + positions = torch.arange( + start_pos, + start_pos + kv.shape[0], + device=mem_index.device, + dtype=torch.long, + ) + infer_state.mem_manager.pack_mla_kv_to_cache( + layer_index=self.layer_num_, + mem_index=mem_index, + kv=kv.reshape(kv.shape[0], 1, kv.shape[-1]), + req_idx=req, + positions=positions, + ) + return + + def _write_compressed_kv(self, infer_state, req, entry_start, comp): + slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, comp.shape[0]) + if comp.shape[0] == 0: + return slots + infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) + return slots + + def _write_c4_indexer_k(self, infer_state, slots, idx_comp): + if idx_comp is None or idx_comp.shape[0] == 0: + return + infer_state.mem_manager.pack_c4_indexer_k_to_cache(self.layer_num_, slots, idx_comp) + return + + def _dense_kv_from_cache(self, infer_state, req, start_pos, end_pos): + if end_pos <= start_pos: + return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") + slots = infer_state.req_manager.req_to_token_indexs[req, start_pos:end_pos].long() + return infer_state.mem_manager.gather_mla_kv(self.layer_num_, slots) + + def _compressed_kv_from_cache(self, infer_state, req, ncomp): + if ncomp == 0: + return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") + if self.compress_ratio == 4: + slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() + else: + slots = infer_state.req_manager.req_to_c128_indexs[req, :ncomp].long() + return infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots) + + def _c4_indexer_k_from_cache(self, infer_state, req, ncomp): + if self.compress_ratio != 4 or ncomp == 0: + return None + slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() + return infer_state.mem_manager.gather_c4_indexer_k(self.layer_num_, slots) + # ------------------------------------------------------------------ attention (prefill) def _attention_prefill(self, x, infer_state, lw): T = x.shape[0] if self.compress_ratio: - cos_tok, sin_tok = infer_state.position_cos_compress, infer_state.position_sin_compress + cos_tok, sin_tok = ( + infer_state.position_cos_compress, + infer_state.position_sin_compress, + ) else: - cos_tok, sin_tok = infer_state.position_cos_sliding, infer_state.position_sin_sliding - q, kv = self._qkv(x, cos_tok, sin_tok, lw) + cos_tok, sin_tok = ( + infer_state.position_cos_sliding, + infer_state.position_sin_sliding, + ) + q, kv, qa = self._qkv(x, cos_tok, sin_tok, lw) sink = lw.attn_sink_.weight o = x.new_empty(T, self.tp_q_heads, self.head_dim) b_req = infer_state.b_req_idx.tolist() starts = infer_state.b_q_start_loc.tolist() lens = infer_state.b_q_seq_len.tolist() - for req, st, ln in zip(b_req, starts, lens): + ready_lens = infer_state.b_ready_cache_len.tolist() + idx_q, idx_weight = self._indexer_q_weight( + x, + qa, + infer_state.position_cos_compress, + infer_state.position_sin_compress, + lw, + ) + for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): q_r, kv_r, x_r = q[st : st + ln], kv[st : st + ln], x[st : st + ln] - kv_all, n_window, ncomp = self._gather_prefill(x_r, kv_r, req, lw, infer_state) - ti = self._topk_idxs_prefill(ln, n_window, ncomp, x.device) - o[st : st + ln] = torch_sparse_attn(q_r.unsqueeze(0), kv_all.unsqueeze(0), sink, ti, self.softmax_scale)[0] + idx_q_r = None if idx_q is None else idx_q[st : st + ln] + idx_weight_r = None if idx_weight is None else idx_weight[st : st + ln] + kv_all, dense_base, n_window, ncomp, idx_comp = self._gather_prefill( + x_r, kv_r, req, ready_len, lw, infer_state + ) + ti = self._topk_idxs_prefill( + ln, + dense_base, + n_window, + ncomp, + x.device, + ready_len, + idx_q_r, + idx_comp, + idx_weight_r, + infer_state, + ) + o[st : st + ln] = vllm_sparse_attn(q_r.unsqueeze(0), kv_all.unsqueeze(0), sink, ti, self.softmax_scale)[0] + self._post_dense_kv( + infer_state, + req, + ready_len, + infer_state.mem_index[st : st + ln], + kv_r, + ) return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) - def _gather_prefill(self, x_r, kv_r, req, lw, infer_state): + def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): ln = kv_r.shape[0] + idx_comp = None + if ready_len > 0: + return self._gather_prefill_extend(x_r, kv_r, req, ready_len, lw, infer_state) if self.compress_ratio: - comp, ks, ss = compressor_prefill_state( + cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) + comp, ks, ss, cstate_pool = compressor_prefill_state( x_r, lw.compressor_wkv_.mm_param.weight, lw.compressor_wgate_.mm_param.weight, @@ -157,27 +267,180 @@ def _gather_prefill(self, x_r, kv_r, req, lw, infer_state): infer_state.cos_compress_table, infer_state.sin_compress_table, self.eps_, + return_state_pool=True, + state_pool=cstate_pool, ) - self._state[req] = {"kv_hist": kv_r.detach(), "comp_kv": comp.detach(), "cstate_kv": ks, "cstate_score": ss} - return torch.cat([kv_r, comp], dim=0), ln, comp.shape[0] - self._state[req] = {"kv_hist": kv_r.detach()} - return kv_r, ln, 0 + comp_slots = self._write_compressed_kv(infer_state, req, 0, comp) + ( + cstate_kv, + cstate_score, + ) = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) + cstate_kv.copy_(ks) + cstate_score.copy_(ss) + state = { + "cstate_kv": cstate_kv, + "cstate_score": cstate_score, + } + if self.compress_ratio == 4: + idx_cstate_pool = infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) + idx_comp, idx_ks, idx_ss, idx_cstate_pool = compressor_prefill_state( + x_r, + lw.idx_cmp_wkv_.mm_param.weight, + lw.idx_cmp_wgate_.mm_param.weight, + lw.idx_cmp_norm_.weight, + lw.idx_cmp_ape_.weight, + 4, + self.index_head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + return_state_pool=True, + state_pool=idx_cstate_pool, + ) + self._write_c4_indexer_k(infer_state, comp_slots, idx_comp) + idx_state = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) + idx_cstate_kv = idx_state[req, 0] + idx_cstate_score = idx_state[req, 1] + idx_cstate_kv.copy_(idx_ks) + idx_cstate_score.copy_(idx_ss) + state.update( + { + "idx_cstate_kv": idx_cstate_kv, + "idx_cstate_score": idx_cstate_score, + } + ) + infer_state.req_manager.set_runtime_state( + req, + self.layer_num_, + state, + ) + ncomp = comp.shape[0] + if DSV4_DEBUG_DISABLE_COMP_ATTN: + return kv_r, 0, ln, 0, None + if not DSV4_DEBUG_DIRECT_PREFILL_COMP: + comp = self._compressed_kv_from_cache(infer_state, req, ncomp) + idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) + return torch.cat([kv_r, comp], dim=0), 0, ln, ncomp, idx_comp + return kv_r, 0, ln, 0, None - def _topk_idxs_prefill(self, seqlen, n_window, ncomp, device): - t = torch.arange(seqlen, device=device) - j = torch.arange(n_window, device=device) - win = torch.where( - (j.unsqueeze(0) <= t.unsqueeze(1)) & (j.unsqueeze(0) > (t.unsqueeze(1) - self.window)), - j.unsqueeze(0).expand(seqlen, n_window), - torch.full((seqlen, n_window), -1, device=device, dtype=torch.long), + def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state): + if self.compress_ratio: + try: + state = infer_state.req_manager.get_runtime_state(req, self.layer_num_) + except KeyError as exc: + raise RuntimeError( + "DeepSeek-V4 prefill chunk is missing runtime state; radix prompt cache " + "must stay disabled until V4 managed token cache is implemented." + ) from exc + cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) + idx_cstate_pool = ( + infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) + if self.compress_ratio == 4 + else None + ) + + for j in range(x_r.shape[0]): + start_pos = ready_len + j + entry = compressor_decode_step( + x_r[j], + lw.compressor_wkv_.mm_param.weight, + lw.compressor_wgate_.mm_param.weight, + lw.compressor_norm_.weight, + lw.compressor_ape_.weight, + self.compress_ratio, + self.head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + state["cstate_kv"], + state["cstate_score"], + start_pos, + state_pool=cstate_pool, + ) + if entry is not None: + entry_start = (start_pos + 1) // self.compress_ratio - 1 + slots = self._write_compressed_kv(infer_state, req, entry_start, entry.unsqueeze(0)) + if self.compress_ratio == 4: + idx_entry = compressor_decode_step( + x_r[j], + lw.idx_cmp_wkv_.mm_param.weight, + lw.idx_cmp_wgate_.mm_param.weight, + lw.idx_cmp_norm_.weight, + lw.idx_cmp_ape_.weight, + 4, + self.index_head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + state["idx_cstate_kv"], + state["idx_cstate_score"], + start_pos, + state_pool=idx_cstate_pool, + ) + if idx_entry is not None: + if entry is None: + entry_start = (start_pos + 1) // self.compress_ratio - 1 + slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, 1) + self._write_c4_indexer_k(infer_state, slots, idx_entry.unsqueeze(0)) + dense_end = ready_len + x_r.shape[0] + ncomp = dense_end // self.compress_ratio + dense_base = max(0, ready_len - self.window + 1) + cached_dense = self._dense_kv_from_cache(infer_state, req, dense_base, ready_len) + dense = torch.cat([cached_dense, kv_r], dim=0) + comp = self._compressed_kv_from_cache(infer_state, req, ncomp) + idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) + if DSV4_DEBUG_DISABLE_COMP_ATTN: + return dense, dense_base, dense.shape[0], 0, None + return ( + torch.cat([dense, comp], dim=0), + dense_base, + dense.shape[0], + ncomp, + idx_comp, + ) + dense_base = max(0, ready_len - self.window + 1) + cached_dense = self._dense_kv_from_cache(infer_state, req, dense_base, ready_len) + dense = torch.cat([cached_dense, kv_r], dim=0) + return ( + dense, + dense_base, + dense.shape[0], + 0, + None, ) + + def _topk_idxs_prefill( + self, + seqlen, + dense_base, + n_window, + ncomp, + device, + base_pos, + idx_q, + idx_comp, + idx_weight, + infer_state, + ): + t = torch.arange(seqlen, device=device) + abs_pos = t + base_pos + offsets = torch.arange(self.window, device=device) + win_abs = abs_pos.unsqueeze(1) - (self.window - 1 - offsets).unsqueeze(0) + valid = (win_abs >= dense_base) & (win_abs < dense_base + n_window) + win = torch.where(valid, win_abs - dense_base, torch.full_like(win_abs, -1)) if ncomp: - c = torch.arange(ncomp, device=device) - comp = torch.where( - c.unsqueeze(0) < ((t.unsqueeze(1) + 1) // self.compress_ratio), - (c.unsqueeze(0) + n_window).expand(seqlen, ncomp), - torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long), - ) + if self.compress_ratio == 4 and ncomp > self.index_topk: + comp = self._indexer_topk(idx_q, idx_comp, idx_weight, abs_pos + 1, n_window, infer_state) + else: + c = torch.arange(ncomp, device=device) + comp = torch.where( + c.unsqueeze(0) < ((abs_pos.unsqueeze(1) + 1) // self.compress_ratio), + (c.unsqueeze(0) + n_window).expand(seqlen, ncomp), + torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long), + ) return torch.cat([win, comp], dim=1).int().unsqueeze(0) return win.int().unsqueeze(0) @@ -185,19 +448,45 @@ def _topk_idxs_prefill(self, seqlen, n_window, ncomp, device): def _attention_decode(self, x, infer_state, lw): B = x.shape[0] # one new token per request if self.compress_ratio: - cos_tok, sin_tok = infer_state.position_cos_compress, infer_state.position_sin_compress + cos_tok, sin_tok = ( + infer_state.position_cos_compress, + infer_state.position_sin_compress, + ) else: - cos_tok, sin_tok = infer_state.position_cos_sliding, infer_state.position_sin_sliding - q, kv = self._qkv(x, cos_tok, sin_tok, lw) # [B, heads, hd], [B, hd] + cos_tok, sin_tok = ( + infer_state.position_cos_sliding, + infer_state.position_sin_sliding, + ) + q, kv, qa = self._qkv(x, cos_tok, sin_tok, lw) # [B, heads, hd], [B, hd] + idx_q, idx_weight = self._indexer_q_weight( + x, + qa, + infer_state.position_cos_compress, + infer_state.position_sin_compress, + lw, + ) sink = lw.attn_sink_.weight b_req = infer_state.b_req_idx.tolist() seqlens = infer_state.b_seq_len.tolist() o = x.new_empty(B, self.tp_q_heads, self.head_dim) for i, (req, seq) in enumerate(zip(b_req, seqlens)): - stt = self._state[req] - stt["kv_hist"] = torch.cat([stt["kv_hist"], kv[i : i + 1]], dim=0) start_pos = seq - 1 + self._post_dense_kv( + infer_state, + req, + start_pos, + infer_state.mem_index[i : i + 1], + kv[i : i + 1], + ) if self.compress_ratio: + try: + stt = infer_state.req_manager.get_runtime_state(req, self.layer_num_) + except KeyError as exc: + raise RuntimeError( + "DeepSeek-V4 decode is missing runtime state; radix prompt cache " + "must stay disabled until V4 managed token cache is implemented." + ) from exc + cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) e = compressor_decode_step( x[i], lw.compressor_wkv_.mm_param.weight, @@ -213,58 +502,257 @@ def _attention_decode(self, x, infer_state, lw): stt["cstate_kv"], stt["cstate_score"], start_pos, + state_pool=cstate_pool, ) + entry_slots = None if e is not None: - stt["comp_kv"] = torch.cat([stt["comp_kv"], e.unsqueeze(0)], dim=0) - win_kv = stt["kv_hist"][-self.window :] - kv_all = torch.cat([win_kv, stt["comp_kv"]], dim=0) + entry_start = (start_pos + 1) // self.compress_ratio - 1 + entry_slots = self._write_compressed_kv(infer_state, req, entry_start, e.unsqueeze(0)) + if self.compress_ratio == 4: + idx_cstate_pool = infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) + idx_e = compressor_decode_step( + x[i], + lw.idx_cmp_wkv_.mm_param.weight, + lw.idx_cmp_wgate_.mm_param.weight, + lw.idx_cmp_norm_.weight, + lw.idx_cmp_ape_.weight, + 4, + self.index_head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + stt["idx_cstate_kv"], + stt["idx_cstate_score"], + start_pos, + state_pool=idx_cstate_pool, + ) + if idx_e is not None: + if entry_slots is None: + entry_start = (start_pos + 1) // self.compress_ratio - 1 + entry_slots = infer_state.req_manager.ensure_compress_slots( + self.layer_num_, req, entry_start, 1 + ) + self._write_c4_indexer_k(infer_state, entry_slots, idx_e.unsqueeze(0)) + win_start = max(0, seq - self.window) + win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) + comp_kv = self._compressed_kv_from_cache(infer_state, req, seq // self.compress_ratio) + idx_comp = self._c4_indexer_k_from_cache(infer_state, req, comp_kv.shape[0]) + if DSV4_DEBUG_DISABLE_COMP_ATTN: + comp_kv = None + idx_comp = None + kv_all = win_kv + else: + kv_all = torch.cat([win_kv, comp_kv], dim=0) else: - win_kv = stt["kv_hist"][-self.window :] + win_start = max(0, seq - self.window) + win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) kv_all = win_kv - ti = torch.arange(kv_all.shape[0], device=x.device).view(1, 1, -1).int() - o[i] = torch_sparse_attn( - q[i].view(1, 1, self.tp_q_heads, self.head_dim), kv_all.unsqueeze(0), sink, ti, self.softmax_scale + comp_kv = None + idx_comp = None + ti = self._topk_idxs_decode( + win_kv.shape[0], + comp_kv, + None if idx_q is None else idx_q[i : i + 1], + idx_comp, + None if idx_weight is None else idx_weight[i : i + 1], + seq, + x.device, + infer_state, + ) + o[i] = vllm_sparse_attn( + q[i].view(1, 1, self.tp_q_heads, self.head_dim), + kv_all.unsqueeze(0), + sink, + ti, + self.softmax_scale, )[0, 0] return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + def _indexer_q_weight(self, x, qa, cos_tok, sin_tok, lw): + if self.compress_ratio != 4: + return None, None + idx_q = lw.idx_wq_b_.mm(qa).view(x.shape[0], self.tp_index_heads, self.index_head_dim) + idx_q = torch.cat( + [ + idx_q[..., : -self.rope_dim], + apply_rotary_emb( + idx_q[..., -self.rope_dim :], + cos_tok.unsqueeze(1), + sin_tok.unsqueeze(1), + ), + ], + dim=-1, + ) + idx_weight = lw.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale + return idx_q, idx_weight + + def _indexer_topk(self, idx_q, idx_comp, idx_weight, positions_1based, offset, infer_state): + ncomp = idx_comp.shape[0] + k = min(self.index_topk, ncomp) + if k == 0: + return torch.empty((idx_q.shape[0], 0), device=idx_q.device, dtype=torch.long) + + scores = torch.einsum("thd,nd->thn", idx_q.float(), idx_comp.float()) + scores = F.relu(scores) * self.indexer_score_scale + index_scores = (scores * idx_weight.unsqueeze(-1)).sum(dim=1) + if self.tp_world_size_ > 1: + all_reduce( + index_scores, + op=dist.ReduceOp.SUM, + group=infer_state.dist_group, + async_op=False, + ) + + causal_threshold = positions_1based // 4 + top = self._indexer_topk_kernel(index_scores, causal_threshold, k) + valid = top >= 0 + return torch.where(valid, top + offset, torch.full_like(top, -1)) + + def _indexer_topk_kernel(self, index_scores, causal_threshold, topk): + if index_scores.is_cuda: + try: + import vllm._C # noqa: F401 + + scores = index_scores.contiguous() + lengths = causal_threshold.to(torch.int32).contiguous() + starts = torch.zeros_like(lengths, dtype=torch.int32) + top = torch.empty((scores.shape[0], topk), dtype=torch.int32, device=scores.device) + torch.ops._C.top_k_per_row_prefill( + scores, + starts, + lengths, + top, + scores.shape[0], + scores.stride(0), + scores.stride(1), + topk, + ) + return top.long() + except Exception: + pass + + entry_indices = torch.arange(index_scores.shape[1], device=index_scores.device) + index_scores = index_scores.masked_fill( + entry_indices.unsqueeze(0) >= causal_threshold.unsqueeze(1), float("-inf") + ) + top = index_scores.topk(topk, dim=-1).indices + valid = top < causal_threshold.unsqueeze(1) + return torch.where(valid, top, torch.full_like(top, -1)) + + def _topk_idxs_decode( + self, + win_len, + comp_kv, + idx_q, + idx_comp, + idx_weight, + seq_len, + device, + infer_state, + ): + win = torch.arange(win_len, device=device, dtype=torch.long) + if comp_kv is None or comp_kv.shape[0] == 0: + return win.view(1, 1, -1).int() + ncomp = comp_kv.shape[0] + if self.compress_ratio == 4 and ncomp > self.index_topk: + comp = self._indexer_topk( + idx_q, + idx_comp, + idx_weight, + torch.tensor([seq_len], device=device, dtype=torch.long), + win_len, + infer_state, + )[0] + else: + comp = torch.arange(ncomp, device=device, dtype=torch.long) + win_len + return torch.cat([win, comp], dim=0).view(1, 1, -1).int() + # ------------------------------------------------------------------ moe def _fp4_experts(self, x, weights, indices, lw): experts = lw.experts_ - out = torch.zeros(x.shape, device=x.device, dtype=torch.float32) - counts = torch.bincount(indices.reshape(-1), minlength=experts.n_routed_experts) - for expert_id in torch.nonzero(counts, as_tuple=False).flatten().tolist(): - token_idx, top_idx = torch.where(indices == expert_id) - if token_idx.numel() == 0: - continue - x_i = x[token_idx] - w1 = dequant_fp4_group_to_bf16(experts.w1[expert_id], experts.w1_scale[expert_id]) - w3 = dequant_fp4_group_to_bf16(experts.w3[expert_id], experts.w3_scale[expert_id]) - gate = F.linear(x_i, w1).float().clamp(max=self.swiglu_limit) - up = F.linear(x_i, w3).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) - hidden = F.silu(gate) * up - hidden.mul_(weights[token_idx, top_idx].unsqueeze(-1)) - w2 = dequant_fp4_group_to_bf16(experts.w2[expert_id], experts.w2_scale[expert_id]) - out.index_add_(0, token_idx, F.linear(hidden.to(x.dtype), w2).float()) - return out.to(x.dtype) + if getattr(experts, "moe_backend", None) != "marlin": + err = getattr(experts, "moe_backend_error", "unknown") + raise RuntimeError(f"DeepSeek-V4 FP4 MoE requires vLLM Marlin backend, init_error={err}") + return self._fp4_experts_marlin(x, weights, indices, experts) + + def _fp4_experts_marlin(self, x, weights, indices, experts): + from vllm.model_executor.layers.fused_moe.activation import MoEActivation + from vllm.model_executor.layers.fused_moe.experts.marlin_moe import ( + fused_marlin_moe, + ) + from vllm.scalar_type import scalar_types + + return fused_marlin_moe( + hidden_states=x.contiguous(), + w1=experts.marlin_w13, + w2=experts.marlin_w2, + bias1=None, + bias2=None, + w1_scale=experts.marlin_w13_scale, + w2_scale=experts.marlin_w2_scale, + topk_weights=weights.to(torch.float32).contiguous(), + topk_ids=indices.to(torch.long).contiguous(), + quant_type_id=scalar_types.float4_e2m1f.id, + global_num_experts=experts.n_routed_experts, + activation=MoEActivation.SILU, + clamp_limit=float(self.swiglu_limit), + ) def _moe_ffn(self, x, infer_state, lw): gw = lw.gate_weight_.mm_param.weight - scores = F.softplus(F.linear(x.float(), gw.float())).sqrt() # sqrtsoftplus - if self.is_hash: - indices = lw.gate_tid2eid_.weight[infer_state.input_ids.long()] - else: - indices = (scores + lw.gate_bias_.weight.unsqueeze(0)).topk(self.topk, dim=-1)[1] - weights = scores.gather(1, indices) - weights = (weights / (weights.sum(-1, keepdim=True) + 1e-20) * self.route_scale).to(torch.float32) - routed = self._fp4_experts(x, weights, indices.long(), lw) + logits = F.linear(x.float(), gw.float()).contiguous() + weights, indices = self._select_experts(logits, infer_state, lw) + routed = self._fp4_experts(x, weights, indices, lw) g = lw.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) u = lw.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) shared = lw.shared_down_.mm((F.silu(g) * u).to(x.dtype)) if self.enable_ep_moe and getattr(lw.experts_, "is_ep", False): if self.tp_world_size_ > 1: - all_reduce(shared, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + all_reduce( + shared, + op=dist.ReduceOp.SUM, + group=infer_state.dist_group, + async_op=False, + ) return routed + shared out = routed + shared if self.tp_world_size_ > 1: all_reduce(out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) return out + + def _select_experts(self, logits, infer_state, lw): + return self._select_experts_vllm(logits, infer_state, lw) + + def _select_experts_vllm(self, logits, infer_state, lw): + from vllm import _custom_ops as ops + + M = logits.shape[0] + bias = None + input_tokens = None + hash_indices_table = None + indices_dtype = torch.int64 + if self.is_hash: + hash_indices_table = lw.gate_tid2eid_.weight + if not hash_indices_table.is_contiguous(): + hash_indices_table = hash_indices_table.contiguous() + indices_dtype = hash_indices_table.dtype + input_tokens = infer_state.input_ids.to(dtype=indices_dtype).contiguous() + else: + bias = lw.gate_bias_.weight + + weights = torch.empty((M, self.topk), dtype=torch.float32, device=logits.device) + indices = torch.empty((M, self.topk), dtype=indices_dtype, device=logits.device) + token_expert_indices = torch.empty((M, self.topk), dtype=torch.int32, device=logits.device) + ops.topk_hash_softplus_sqrt( + weights, + indices, + token_expert_indices, + logits, + True, + self.route_scale, + bias, + input_tokens, + hash_indices_table, + ) + return weights, indices.long() diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 7c12f714db..cdaaac2cdb 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -1,3 +1,5 @@ +import threading + import torch from lightllm.common.basemodel import TransformerLayerWeight from lightllm.common.basemodel.layer_weights.meta_weights import ( @@ -10,10 +12,16 @@ ) from lightllm.common.basemodel.layer_weights.meta_weights.base_weight import BaseWeightTpl from lightllm.common.quantization.registry import QUANTMETHODS +from lightllm.utils.log_utils import init_logger from ..triton_kernel.quant_convert import dequant_fp8_block_to_bf16 +logger = init_logger(__name__) + + class DeepseekV4FP4ExpertsWeight(BaseWeightTpl): + _marlin_pack_lock = threading.Lock() + def __init__(self, weight_prefix, n_routed_experts, hidden_size, moe_intermediate_size, data_type): super().__init__(data_type=data_type) self.weight_prefix = weight_prefix @@ -23,10 +31,21 @@ def __init__(self, weight_prefix, n_routed_experts, hidden_size, moe_intermediat self.split_inter_size = moe_intermediate_size // self.tp_world_size_ self.local_expert_ids = list(range(n_routed_experts)) self.expert_idx_to_local_idx = {expert_idx: expert_idx for expert_idx in self.local_expert_ids} - self._create_weight() + self.moe_backend = None + self.moe_backend_error = None + self._marlin_checked = False + self._load_lock = threading.Lock() + self.load_ok = { + name: [False] * n_routed_experts for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale") + } def _create_weight(self): - device = f"cuda:{self.device_id_}" + self._ensure_raw_fp4_weight() + + def _ensure_raw_fp4_weight(self): + if hasattr(self, "w1"): + return + device = "cpu" n = self.n_routed_experts h = self.hidden_size inter = self.split_inter_size @@ -36,10 +55,6 @@ def _create_weight(self): self.w1_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) self.w3_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) self.w2_scale = torch.empty((n, h, inter // 32), dtype=torch.float8_e8m0fnu, device=device) - self.load_ok = { - name: [False] * n - for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale") - } def _copy_expert_weight(self, dst, weight, expert_idx, name, is_down=False): if is_down: @@ -66,30 +81,122 @@ def _copy_expert_scale(self, dst, scale, expert_idx, name, is_down=False): self.load_ok[name][expert_idx] = True def load_hf_weights(self, weights): + if self._marlin_checked: + return + has_weight = False for expert_idx in self.local_expert_ids: prefix = f"{self.weight_prefix}.{expert_idx}" - w1 = f"{prefix}.w1.weight" - w1_scale = f"{prefix}.w1.scale" - w2 = f"{prefix}.w2.weight" - w2_scale = f"{prefix}.w2.scale" - w3 = f"{prefix}.w3.weight" - w3_scale = f"{prefix}.w3.scale" - if w1 in weights: - self._copy_expert_weight(self.w1, weights[w1], expert_idx, "w1") - if w1_scale in weights: - self._copy_expert_scale(self.w1_scale, weights[w1_scale], expert_idx, "w1_scale") - if w3 in weights: - self._copy_expert_weight(self.w3, weights[w3], expert_idx, "w3") - if w3_scale in weights: - self._copy_expert_scale(self.w3_scale, weights[w3_scale], expert_idx, "w3_scale") - if w2 in weights: - self._copy_expert_weight(self.w2, weights[w2], expert_idx, "w2", is_down=True) - if w2_scale in weights: - self._copy_expert_scale(self.w2_scale, weights[w2_scale], expert_idx, "w2_scale", is_down=True) + if ( + f"{prefix}.w1.weight" in weights + or f"{prefix}.w1.scale" in weights + or f"{prefix}.w2.weight" in weights + or f"{prefix}.w2.scale" in weights + or f"{prefix}.w3.weight" in weights + or f"{prefix}.w3.scale" in weights + ): + has_weight = True + break + if not has_weight: + return + + with self._load_lock: + if self._marlin_checked: + return + self._ensure_raw_fp4_weight() + for expert_idx in self.local_expert_ids: + prefix = f"{self.weight_prefix}.{expert_idx}" + w1 = f"{prefix}.w1.weight" + w1_scale = f"{prefix}.w1.scale" + w2 = f"{prefix}.w2.weight" + w2_scale = f"{prefix}.w2.scale" + w3 = f"{prefix}.w3.weight" + w3_scale = f"{prefix}.w3.scale" + if w1 in weights: + self._copy_expert_weight(self.w1, weights[w1], expert_idx, "w1") + if w1_scale in weights: + self._copy_expert_scale(self.w1_scale, weights[w1_scale], expert_idx, "w1_scale") + if w3 in weights: + self._copy_expert_weight(self.w3, weights[w3], expert_idx, "w3") + if w3_scale in weights: + self._copy_expert_scale(self.w3_scale, weights[w3_scale], expert_idx, "w3_scale") + if w2 in weights: + self._copy_expert_weight(self.w2, weights[w2], expert_idx, "w2", is_down=True) + if w2_scale in weights: + self._copy_expert_scale(self.w2_scale, weights[w2_scale], expert_idx, "w2_scale", is_down=True) + if self._raw_load_complete(): + self._try_init_marlin() def verify_load(self): + with self._load_lock: + ok = self._raw_load_complete() + if ok and not self._marlin_checked: + self._try_init_marlin() + return ok + + def _raw_load_complete(self): return all(all(ok_list) for ok_list in self.load_ok.values()) + def _try_init_marlin(self): + try: + from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( + prepare_moe_mxfp4_layer_for_marlin, + ) + + class _MarlinLayer: + pass + + with self._marlin_pack_lock: + torch.cuda.set_device(self.device_id_) + device = torch.device("cuda", self.device_id_) + layer = _MarlinLayer() + layer.params_dtype = self.data_type_ + w13_cpu, w13_scale_cpu = self._build_w13_weight() + w13 = w13_cpu.to(device=device, non_blocking=True).contiguous() + w2 = self.w2.view(torch.uint8).to(device=device, non_blocking=True).contiguous() + w13_scale = w13_scale_cpu.to(device=device, non_blocking=True).contiguous() + w2_scale = self.w2_scale.to(device=device, non_blocking=True).contiguous() + ( + self.marlin_w13, + self.marlin_w2, + self.marlin_w13_scale, + self.marlin_w2_scale, + _, + _, + ) = prepare_moe_mxfp4_layer_for_marlin(layer, w13, w2, w13_scale, w2_scale, None, None) + del w13_cpu, w13_scale_cpu, w13, w2, w13_scale, w2_scale + self.moe_backend = "marlin" + self._marlin_checked = True + self._release_raw_fp4_weight() + torch.cuda.empty_cache() + logger.info( + "DeepSeek-V4 FP4 experts use vLLM Marlin backend, prefix=%s, rank=%s", + self.weight_prefix, + self.tp_rank_, + ) + except Exception as e: + self.moe_backend_error = repr(e) + raise RuntimeError( + "DeepSeek-V4 FP4 experts require vLLM Marlin backend, " + f"prefix={self.weight_prefix}, rank={self.tp_rank_}, error={self.moe_backend_error}" + ) from e + + def _build_w13_weight(self): + n = self.n_routed_experts + h = self.hidden_size + inter = self.split_inter_size + w13 = torch.empty((n, 2 * inter, h // 2), dtype=torch.uint8, device=self.w1.device) + w13[:, :inter, :].copy_(self.w1.view(torch.uint8)) + w13[:, inter:, :].copy_(self.w3.view(torch.uint8)) + w13_scale = torch.empty((n, 2 * inter, h // 32), dtype=self.w1_scale.dtype, device=self.w1_scale.device) + w13_scale[:, :inter, :].copy_(self.w1_scale) + w13_scale[:, inter:, :].copy_(self.w3_scale) + return w13.contiguous(), w13_scale.contiguous() + + def _release_raw_fp4_weight(self): + for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale"): + if hasattr(self, name): + delattr(self, name) + class DeepseekV4TransformerLayerWeight(TransformerLayerWeight): """Per-layer weights for DeepSeek-V4-Flash. diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 02c71f01b2..c87f2fdebd 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -1,16 +1,37 @@ +import importlib.util +import os + import torch from lightllm.models.registry import ModelRegistry from lightllm.models.llama.model import LlamaTpPartModel -from lightllm.common.req_manager import ReqManager, DeepseekV4ReqManager +from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.basemodel.attention.triton.fp import TritonAttBackend -from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import DeepseekV4PreAndPostLayerWeight -from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight -from lightllm.models.deepseek_v4.layer_infer.pre_layer_infer import DeepseekV4PreLayerInfer -from lightllm.models.deepseek_v4.layer_infer.post_layer_infer import DeepseekV4PostLayerInfer -from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer +from lightllm.common.basemodel.attention.base_att import ( + BaseAttBackend, + BasePrefillAttState, + BaseDecodeAttState, +) +from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4PreAndPostLayerWeight, +) +from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import ( + DeepseekV4TransformerLayerWeight, +) +from lightllm.models.deepseek_v4.layer_infer.pre_layer_infer import ( + DeepseekV4PreLayerInfer, +) +from lightllm.models.deepseek_v4.layer_infer.post_layer_infer import ( + DeepseekV4PostLayerInfer, +) +from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import ( + DeepseekV4TransformerLayerInfer, +) from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo -from lightllm.models.llama.yarn_rotary_utils import find_correction_range, linear_ramp_mask +from lightllm.models.deepseek3_2.model import DeepSeekChatTokenizerBase +from lightllm.models.llama.yarn_rotary_utils import ( + find_correction_range, + linear_ramp_mask, +) from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager @@ -18,9 +39,38 @@ logger = init_logger(__name__) +class DeepseekV4DirectSparseAttBackend(BaseAttBackend): + """Lifecycle placeholder for V4 direct attention. + + V4 attention is currently driven inside the layer by `vllm_sparse_attn()`, not by the generic + `infer_state.prefill_att_state.prefill_att()` / `decode_att()` backend selector. + """ + + def create_att_prefill_state(self, infer_state): + return DeepseekV4DirectSparsePrefillAttState(backend=self, infer_state=infer_state) + + def create_att_decode_state(self, infer_state): + return DeepseekV4DirectSparseDecodeAttState(backend=self, infer_state=infer_state) + + +class DeepseekV4DirectSparsePrefillAttState(BasePrefillAttState): + def init_state(self): + return + + def prefill_att(self, *args, **kwargs): + raise RuntimeError("DeepSeek-V4 attention is executed directly by vllm_sparse_attn() in layer_infer.") + + +class DeepseekV4DirectSparseDecodeAttState(BaseDecodeAttState): + def init_state(self): + return + + def decode_att(self, *args, **kwargs): + raise RuntimeError("DeepSeek-V4 attention is executed directly by vllm_sparse_attn() in layer_infer.") + + @ModelRegistry("deepseek_v4") class DeepseekV4TpPartModel(LlamaTpPartModel): - pre_and_post_weight_class = DeepseekV4PreAndPostLayerWeight transformer_weight_class = DeepseekV4TransformerLayerWeight @@ -50,35 +100,46 @@ def _init_req_manager(self): create_max_seq_len = max(create_max_seq_len, self.max_seq_length) self._dsv4_req_manager_seq_len = create_max_seq_len - self.req_manager = ReqManager(self.max_req_num, create_max_seq_len, None) + layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() + self._dsv4_compress_rates = self._get_compress_rates(layer_num) + self.req_manager = DeepseekV4ReqManager( + self.max_req_num, + create_max_seq_len, + compress_rates=self._dsv4_compress_rates, + head_dim=self.config["head_dim"], + indexer_head_dim=self.config["index_head_dim"], + ) return def _get_compress_rates(self, layer_num): - rates = list(self.config.get("compress_ratios", [])) - if len(rates) < layer_num: - rates.extend([0] * (layer_num - len(rates))) + rates = list(self.config["compress_ratios"]) + assert ( + len(rates) >= layer_num + ), f"DeepSeek-V4 compress_ratios length {len(rates)} is shorter than layer_num {layer_num}" return rates[:layer_num] def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() + compress_rates = getattr(self, "_dsv4_compress_rates", self._get_compress_rates(layer_num)) self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, head_num=1, head_dim=self.config["head_dim"], layer_num=layer_num, - compress_rates=self._get_compress_rates(layer_num), + compress_rates=compress_rates, indexer_head_dim=self.config["index_head_dim"], + max_request_num=self.max_req_num, + sliding_window=self.config["sliding_window"], mem_fraction=self.mem_fraction, ) - self.req_manager = DeepseekV4ReqManager( - self.max_req_num, self._dsv4_req_manager_seq_len, self.mem_manager - ) + assert isinstance(self.req_manager, DeepseekV4ReqManager) + self.req_manager.bind_mem_manager(self.mem_manager) return def _init_att_backend(self): - self.prefill_att_backend = TritonAttBackend(model=self) - self.decode_att_backend = TritonAttBackend(model=self) + self.prefill_att_backend = DeepseekV4DirectSparseAttBackend(model=self) + self.decode_att_backend = DeepseekV4DirectSparseAttBackend(model=self) return def _init_custom(self): @@ -114,8 +175,49 @@ def build(base, factor, orig_max): f = torch.outer(torch.arange(max_seq, dtype=torch.float32, device="cuda"), freqs) # [max_seq, dim//2] return f.cos(), f.sin() - self._cos_cached_sliding, self._sin_cached_sliding = build(cfg["rope_theta"], rs.get("factor", 1.0), 0) + self._cos_cached_sliding, self._sin_cached_sliding = build( + cfg["rope_theta"], + rs.get("factor", 16), + rs.get("original_max_position_embeddings", 65536), + ) self._cos_cached_compress, self._sin_cached_compress = build( - cfg["compress_rope_theta"], rs.get("factor", 16), rs.get("original_max_position_embeddings", 65536) + cfg["compress_rope_theta"], + rs.get("factor", 16), + rs.get("original_max_position_embeddings", 65536), ) return + + +class DeepSeekV4Tokenizer(DeepSeekChatTokenizerBase): + """Tokenizer wrapper for DeepSeek-V4's Python prompt encoding.""" + + def __init__(self, tokenizer, model_dir): + super().__init__(tokenizer) + self.model_dir = model_dir + self._encoding_module = None + + def _get_encoding_module(self): + if self._encoding_module is not None: + return self._encoding_module + + encoding_path = os.path.join(self.model_dir, "encoding", "encoding_dsv4.py") + if not os.path.exists(encoding_path): + raise FileNotFoundError(f"DeepSeek-V4 encoding file not found: {encoding_path}") + + spec = importlib.util.spec_from_file_location("lightllm_deepseek_v4_encoding_dsv4", encoding_path) + if spec is None or spec.loader is None: + raise ImportError(f"failed to load DeepSeek-V4 encoding module from {encoding_path}") + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + self._encoding_module = module + return module + + def _encode_messages(self, msgs, thinking_mode, kwargs): + encoding = self._get_encoding_module() + return encoding.encode_messages( + msgs, + thinking_mode=thinking_mode, + drop_thinking=kwargs.get("drop_thinking", True), + add_default_bos_token=kwargs.get("add_default_bos_token", True), + reasoning_effort=kwargs.get("reasoning_effort"), + ) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 654ba0f3e5..5a92a339cb 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -10,7 +10,11 @@ from .metrics.manager import start_metric_manager from .embed_cache.manager import start_cache_manager from lightllm.utils.log_utils import init_logger -from lightllm.utils.envs_utils import set_env_start_args, set_unique_server_name, get_unique_server_name +from lightllm.utils.envs_utils import ( + set_env_start_args, + set_unique_server_name, + get_unique_server_name, +) from lightllm.utils.envs_utils import get_lightllm_gunicorn_keep_alive from .detokenization.manager import start_detokenization_process from .router.manager import start_router_process @@ -23,6 +27,8 @@ has_vision_module, is_linear_att_mixed_model, auto_set_max_req_total_len, + get_model_type, + get_config_json, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args @@ -83,7 +89,14 @@ def normal_or_p_d_start(args): enable_mps() - if args.run_mode not in ["normal", "prefill", "decode", "nixl_prefill", "nixl_decode", "visual_only"]: + if args.run_mode not in [ + "normal", + "prefill", + "decode", + "nixl_prefill", + "nixl_decode", + "visual_only", + ]: return # 通过模型的参数判断是否是多模态模型,包含哪几种模态, 并设置是否启动相应得模块 @@ -108,6 +121,23 @@ def normal_or_p_d_start(args): else: args.enable_multimodal = True + model_type = get_model_type(args.model_dir) + if model_type == "deepseek_v4": + if args.run_mode != "normal": + raise NotImplementedError("DeepSeek-V4 currently supports only run_mode=normal in LightLLM.") + if args.enable_cpu_cache or args.enable_disk_cache: + raise NotImplementedError("DeepSeek-V4 CPU/disk KV cache is not supported yet.") + if args.mtp_mode is not None or args.mtp_draft_model_dir is not None or args.mtp_step != 0: + raise NotImplementedError("DeepSeek-V4 MTP/speculative decoding is not supported yet.") + if args.enable_ep_moe: + raise NotImplementedError("DeepSeek-V4 EP MoE is not supported yet; use TP for now.") + if "prompt_cache_kv_buffer" in get_config_json(args.model_dir): + raise NotImplementedError("DeepSeek-V4 prompt_cache_kv_buffer is not supported yet.") + if not args.disable_dynamic_prompt_cache: + logger.info("DeepSeek-V4 runtime state does not support radix prompt cache yet; disabling it.") + args.disable_dynamic_prompt_cache = True + args.use_dynamic_prompt_cache = False + if args.enable_cpu_cache: # 生成一个用于创建cpu kv cache的共享内存id。 args.cpu_kv_cache_shm_id = uuid.uuid1().int % 123456789 @@ -333,7 +363,14 @@ def normal_or_p_d_start(args): from lightllm.utils.config_utils import get_dtype args.data_type = get_dtype(args.model_dir) - assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] + assert args.data_type in [ + "fp16", + "float16", + "bf16", + "bfloat16", + "fp32", + "float32", + ] already_uesd_ports = [args.port] if args.nccl_port is not None: @@ -432,7 +469,6 @@ def normal_or_p_d_start(args): ) if not args.disable_vision: - if not args.visual_use_proxy_mode: from .visualserver.manager import start_visual_process @@ -616,7 +652,14 @@ def visual_only_start(args): from lightllm.utils.config_utils import get_dtype args.data_type = get_dtype(args.model_dir) - assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] + assert args.data_type in [ + "fp16", + "float16", + "bf16", + "bfloat16", + "fp32", + "float32", + ] logger.info(f"alloced ports: {can_use_ports}") diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 5e90c9b34a..abeb8d61e9 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -51,7 +51,9 @@ def register( vocab_size: int, ): self.args = get_env_start_args() - from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend + from lightllm.server.router.model_infer.mode_backend.base_backend import ( + ModeBackend, + ) self.backend: ModeBackend = backend self.req_manager = req_manager @@ -122,7 +124,21 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: return req_objs - def free_a_req_mem(self, free_token_index: List, req: "InferReq"): + def free_a_req_mem( + self, + free_token_index: List, + req: "InferReq", + free_c4_index: Optional[List] = None, + free_c128_index: Optional[List] = None, + ): + if hasattr(self.req_manager, "pop_compress_indices_for_req"): + c4, c128 = self.req_manager.pop_compress_indices_for_req(req.req_idx) + if c4 is not None and free_c4_index is not None: + free_c4_index.append(c4) + if c128 is not None and free_c128_index is not None: + free_c128_index.append(c128) + self.req_manager.clear_runtime_state(req.req_idx) + if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) else: @@ -258,11 +274,13 @@ def _filter(self, finished_request_ids: List[int]): free_req_index = [] free_token_index = [] + free_c4_index = [] + free_c128_index = [] for request_id in finished_request_ids: req: InferReq = self.requests_mapping.pop(request_id) if self.args.diverse_mode: req.clear_master_slave_state() - self.free_a_req_mem(free_token_index, req) + self.free_a_req_mem(free_token_index, req, free_c4_index, free_c128_index) free_req_index.append(req.req_idx) # logger.info(f"infer release req id {req.shm_req.request_id}") @@ -270,7 +288,17 @@ def _filter(self, finished_request_ids: List[int]): self.shm_req_manager.put_back_req_obj(req.shm_req) free_token_index = custom_cat(free_token_index) - self.req_manager.free(free_req_index, free_token_index) + if hasattr(self.req_manager, "free_compress_indices"): + free_c4_index = custom_cat(free_c4_index) if free_c4_index else None + free_c128_index = custom_cat(free_c128_index) if free_c128_index else None + self.req_manager.free( + free_req_index, + free_token_index, + free_c4_index=free_c4_index, + free_c128_index=free_c128_index, + ) + else: + self.req_manager.free(free_req_index, free_token_index) finished_req_ids_set = set(finished_request_ids) self.infer_req_ids = [_id for _id in self.infer_req_ids if _id not in finished_req_ids_set] @@ -299,11 +327,13 @@ def pause_reqs(self, pause_reqs: List["InferReq"], is_master_in_dp: bool): g_infer_state_lock.acquire() free_token_index = [] + free_c4_index = [] + free_c128_index = [] for req in pause_reqs: if self.args.diverse_mode: # 发生暂停的时候,需要清除 diverse 模式下的主从关系 req.clear_master_slave_state() - self.free_a_req_mem(free_token_index, req) + self.free_a_req_mem(free_token_index, req, free_c4_index, free_c128_index) assert req.wait_pause is True req.wait_pause = False req.paused = True @@ -314,11 +344,23 @@ def pause_reqs(self, pause_reqs: List["InferReq"], is_master_in_dp: bool): if len(free_token_index) != 0: free_token_index = custom_cat(free_token_index) self.req_manager.free_token(free_token_index) + if hasattr(self.req_manager, "free_compress_indices"): + free_c4_index = custom_cat(free_c4_index) if free_c4_index else None + free_c128_index = custom_cat(free_c128_index) if free_c128_index else None + self.req_manager.free_compress_indices( + free_c4_index=free_c4_index, + free_c128_index=free_c128_index, + ) g_infer_state_lock.release() return self - def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bool, can_alloc_token_num: int): + def recover_paused_reqs( + self, + paused_reqs: List["InferReq"], + is_master_in_dp: bool, + can_alloc_token_num: int, + ): if paused_reqs: g_infer_state_lock.acquire() @@ -375,7 +417,9 @@ def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: L big_page_buffer_ids = torch.tensor(big_page_buffer_ids, dtype=torch.int32, requires_grad=False, device="cpu") big_page_buffer_ids = big_page_buffer_ids.cuda(non_blocking=True) - from lightllm.common.basemodel.triton_kernel.linear_att_copy import copy_linear_att_state_to_kv_buffer + from lightllm.common.basemodel.triton_kernel.linear_att_copy import ( + copy_linear_att_state_to_kv_buffer, + ) copy_linear_att_state_to_kv_buffer( b_req_idx=b_req_idx, @@ -405,9 +449,10 @@ def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: L gpu_ssm_state = self.req_manager.req_to_ssm_state.buffer[:, src_buffer_idx, ...] dst_buffer_idx = req.tail_linear_att_small_page_buffer_id - dst_conv_state, dst_ssm_state = self.radix_cache.linear_att_small_page_buffers.get_state_cache( - buffer_idx=dst_buffer_idx - ) + ( + dst_conv_state, + dst_ssm_state, + ) = self.radix_cache.linear_att_small_page_buffers.get_state_cache(buffer_idx=dst_buffer_idx) # TODO 对于非连续对象调用 copy_ 效率并不高 dst_conv_state.copy_(gpu_conv_state, non_blocking=True) dst_ssm_state.copy_(gpu_ssm_state, non_blocking=True) @@ -640,7 +685,10 @@ def _linear_match_radix_cache(self): enable_prompt_cache = (not self.sampling_param.disable_prompt_cache) and g_infer_context.radix_cache is not None linear_hash_list = self.shm_req.linear_att_token_hash_list.get_all() linear_att_hash_page_size = self.args.linear_att_hash_page_size - match_tokens = min(len(linear_hash_list) * linear_att_hash_page_size, self.get_cur_total_len() - 1) + match_tokens = min( + len(linear_hash_list) * linear_att_hash_page_size, + self.get_cur_total_len() - 1, + ) match_tokens = max(0, match_tokens) match_tokens = (match_tokens // linear_att_hash_page_size) * linear_att_hash_page_size match_block_num = match_tokens // linear_att_hash_page_size @@ -706,7 +754,8 @@ def _linear_match_radix_cache(self): # 将 对应的 value_tensors 中的 kv 数据 拷贝到 tail_mems 中对应的数据去 radix_cache.mem_manager.operator.copy_mem_to_mem( - value_tensor[cur_big_page_tokens:shared_kv_len], tail_mems + value_tensor[cur_big_page_tokens:shared_kv_len], + tail_mems, ) self.shared_kv_node = share_node # 只是为了保证 copy_small_page_buffer_to_linear_att_state 正确调用 @@ -737,7 +786,8 @@ def _linear_match_radix_cache(self): assert self.tail_linear_att_small_page_buffer_id is None # 恢复linear att 状态 g_infer_context.req_manager.copy_big_page_buffer_to_linear_att_state( - big_page_buffer_idx=share_node.big_page_buffer_idx, req=self + big_page_buffer_idx=share_node.big_page_buffer_idx, + req=self, ) self.shm_req.shm_cur_kv_len = self.cur_kv_len diff --git a/lightllm/server/tokenizer.py b/lightllm/server/tokenizer.py index f84e6359ba..18d5eafcc0 100644 --- a/lightllm/server/tokenizer.py +++ b/lightllm/server/tokenizer.py @@ -89,6 +89,11 @@ def get_tokenizer( ) logger.info("Using DeepSeek-V3.2 tokenizer mode with Python-based chat template encoding.") return DeepSeekV32Tokenizer(hf_tokenizer) + if model_type == "deepseek_v4": + from ..models.deepseek_v4.model import DeepSeekV4Tokenizer + + logger.info("Using DeepSeek-V4 tokenizer mode with Python-based chat template encoding.") + return DeepSeekV4Tokenizer(tokenizer, tokenizer_name) if model_cfg["architectures"][0] == "TarsierForConditionalGeneration": from ..models.qwen2_vl.vision_process import Qwen2VLImageProcessor From a1612445898361b4a48a6f4c4778ad2043b24072 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 5 Jun 2026 05:03:58 +0000 Subject: [PATCH 003/214] add prompt cache --- lightllm/common/basemodel/basemodel.py | 28 ++ .../deepseek4_mem_manager.py | 204 +++++++++++++- lightllm/common/req_manager.py | 255 ++++++++++++++++++ .../layer_infer/transformer_layer_infer.py | 32 ++- lightllm/server/api_start.py | 4 - .../router/dynamic_prompt/radix_cache.py | 172 +++++++++--- .../server/router/model_infer/infer_batch.py | 124 ++++++++- .../model_infer/mode_backend/base_backend.py | 38 +++ 8 files changed, 803 insertions(+), 54 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 473dcbafda..d785991808 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -527,12 +527,22 @@ def _prefill( alloc_mem_index=infer_state.mem_index, max_q_seq_len=infer_state.max_q_seq_len, ) + if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): + self.mem_manager.prepare_prefill_swa_slots( + b_req_idx=infer_state.b_req_idx, + b_seq_len=infer_state.b_seq_len, + b_ready_cache_len=infer_state.b_ready_cache_len, + b_start_loc=model_input.b_prefill_start_loc, + mem_index=infer_state.mem_index, + ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() infer_state.init_some_extra_state(self) infer_state.init_att_state() model_output = self._context_forward(infer_state) + if hasattr(self.mem_manager, "commit_prefill_swa_slots"): + self.mem_manager.commit_prefill_swa_slots() model_output = self._create_unpad_prefill_model_output( padded_model_output=model_output, @@ -747,6 +757,14 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state0.mem_index, max_q_seq_len=infer_state0.max_q_seq_len, ) + if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): + self.mem_manager.prepare_prefill_swa_slots( + b_req_idx=infer_state0.b_req_idx, + b_seq_len=infer_state0.b_seq_len, + b_ready_cache_len=infer_state0.b_ready_cache_len, + b_start_loc=model_input0.b_prefill_start_loc, + mem_index=infer_state0.mem_index, + ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -760,6 +778,14 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state1.mem_index, max_q_seq_len=infer_state1.max_q_seq_len, ) + if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): + self.mem_manager.prepare_prefill_swa_slots( + b_req_idx=infer_state1.b_req_idx, + b_seq_len=infer_state1.b_seq_len, + b_ready_cache_len=infer_state1.b_ready_cache_len, + b_start_loc=model_input1.b_prefill_start_loc, + mem_index=infer_state1.mem_index, + ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -767,6 +793,8 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod prefill_mem_indexes_ready_event.record() model_output0, model_output1 = self._overlap_tpsp_context_forward(infer_state0, infer_state1=infer_state1) + if hasattr(self.mem_manager, "commit_prefill_swa_slots"): + self.mem_manager.commit_prefill_swa_slots() model_output0 = self._create_unpad_prefill_model_output( padded_model_output=model_output0, diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 739e0bd51a..6735b2deed 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -28,7 +28,7 @@ DSV4_SWA_PAGE_SIZE = 128 DSV4_C4_PAGE_SIZE = 64 DSV4_C128_PAGE_SIZE = 2 -DSV4_PROFILE_MAX_FULL_TOKENS = 2_000_000 +DSV4_PROFILE_MAX_FULL_TOKENS = 1_500_000 def _ceil_div(a: int, b: int) -> int: @@ -294,6 +294,7 @@ def __init__( self.cache_dtype = torch.uint8 self.max_request_num = max_request_num self.sliding_window = sliding_window + self._pending_prefill_swa: Dict[int, Dict[str, torch.Tensor]] = {} # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) self.layer_to_c4_idx: Dict[int, int] = {} @@ -560,6 +561,127 @@ def ensure_swa_slots(self, req_idx: int, positions: torch.Tensor, full_slots: to out[i] = swa return out + def _reserve_prefill_swa_slots( + self, + req_idx: int, + positions: torch.Tensor, + full_slots: torch.Tensor, + ) -> Dict[str, torch.Tensor]: + full_slots = full_slots.long().reshape(-1) + positions = positions.long().reshape(-1) + assert positions.numel() == full_slots.numel() + + out = torch.empty_like(full_slots, dtype=torch.long) + ring_to_swa: Dict[int, int] = {} + ring_to_old_full: Dict[int, int] = {} + ring_to_final_full: Dict[int, int] = {} + hold = self.swa_pool.HOLD_TOKEN_MEMINDEX + + for i, (pos, full) in enumerate(zip(positions.tolist(), full_slots.tolist())): + if full == self.HOLD_TOKEN_MEMINDEX: + out[i] = hold + continue + + ring_pos = int(pos) % int(self.sliding_window) + swa = ring_to_swa.get(ring_pos) + if swa is None: + old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) + old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) + if old_swa == hold: + old_swa = int(self.swa_allocator.alloc(1)[0].item()) + swa = old_swa + ring_to_swa[ring_pos] = swa + ring_to_old_full[ring_pos] = old_full + + ring_to_final_full[ring_pos] = int(full) + out[i] = swa + + rings = sorted(ring_to_final_full) + return { + "positions": positions.detach().clone(), + "full_slots": full_slots.detach().clone(), + "swa_slots": out.detach().clone(), + "commit_rings": torch.tensor(rings, dtype=torch.long, device=full_slots.device), + "commit_full_slots": torch.tensor( + [ring_to_final_full[r] for r in rings], + dtype=torch.long, + device=full_slots.device, + ), + "commit_swa_slots": torch.tensor( + [ring_to_swa[r] for r in rings], + dtype=torch.long, + device=full_slots.device, + ), + "commit_old_full_slots": torch.tensor( + [ring_to_old_full[r] for r in rings], + dtype=torch.long, + device=full_slots.device, + ), + } + + def prepare_prefill_swa_slots( + self, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_start_loc: torch.Tensor, + mem_index: torch.Tensor, + ) -> None: + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: + return + + self._pending_prefill_swa = {} + req_list = b_req_idx.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + ready_list = b_ready_cache_len.detach().cpu().tolist() + start_list = b_start_loc.detach().cpu().tolist() + for req_idx, seq_len, ready_len, start_loc in zip(req_list, seq_list, ready_list, start_list): + token_num = int(seq_len) - int(ready_len) + if token_num <= 0: + continue + pos = torch.arange(int(ready_len), int(seq_len), dtype=torch.long, device=mem_index.device) + slots = mem_index[int(start_loc) : int(start_loc) + token_num] + self._pending_prefill_swa[int(req_idx)] = self._reserve_prefill_swa_slots(int(req_idx), pos, slots) + return + + def _get_pending_prefill_swa_slots( + self, + req_idx: int, + positions: torch.Tensor, + full_slots: torch.Tensor, + ) -> Optional[torch.Tensor]: + pending = self._pending_prefill_swa.get(int(req_idx)) + if pending is None: + return None + if pending["positions"].numel() != positions.numel(): + return None + if not torch.equal(pending["positions"].to(positions.device), positions.long().reshape(-1)): + return None + if not torch.equal(pending["full_slots"].to(full_slots.device), full_slots.long().reshape(-1)): + return None + return pending["swa_slots"].to(full_slots.device) + + def commit_prefill_swa_slots(self) -> None: + if not self._pending_prefill_swa: + return + for req_idx, pending in self._pending_prefill_swa.items(): + rings = pending["commit_rings"].to(self.req_to_swa_indexs.device) + if rings.numel() == 0: + continue + old_full = pending["commit_old_full_slots"].to(self.full_to_swa_indexs.device) + valid_old = old_full >= 0 + if valid_old.any(): + self.full_to_swa_indexs[old_full[valid_old].long()] = -1 + + full_slots = pending["commit_full_slots"].to(self.full_to_swa_indexs.device) + swa_slots = pending["commit_swa_slots"].to(self.full_to_swa_indexs.device) + self.req_to_swa_indexs[int(req_idx), rings] = swa_slots.to(torch.int32) + self.req_to_swa_full_indexs[int(req_idx), rings] = full_slots.to(torch.int32) + self.full_to_swa_indexs[full_slots.long()] = swa_slots.to(torch.int32) + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + self._pending_prefill_swa = {} + return + def _swa_slots_from_full(self, full_slots: torch.Tensor) -> torch.Tensor: full_slots = full_slots.long().reshape(-1) if full_slots.numel() == 0: @@ -597,6 +719,79 @@ def free_swa_for_req(self, req_idx: int) -> None: self.req_to_swa_full_indexs[req_idx].fill_(-1) self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + def snapshot_swa_for_prompt_cache(self, req_idx: int, cache_len: int, full_slots: torch.Tensor): + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None or cache_len <= 0: + return None + tail_start = max(0, int(cache_len) - int(self.sliding_window)) + full_slots = full_slots[tail_start:cache_len].long().to(self.kv_buffer.device) + if full_slots.numel() == 0: + return None + swa_slots = self.full_to_swa_indexs[full_slots].long() + if (swa_slots < 0).any(): + bad = int(full_slots[swa_slots < 0][0].item()) + raise RuntimeError(f"DeepSeek-V4 prompt cache cannot snapshot evicted SWA full slot {bad}") + return { + "positions": torch.arange(tail_start, cache_len, dtype=torch.int64, device="cpu"), + "full_slots": full_slots.detach().cpu(), + "swa_slots": swa_slots.detach().cpu(), + } + + def clone_swa_for_prompt_cache(self, req_idx: int, cache_len: int, full_slots: torch.Tensor): + payload = self.snapshot_swa_for_prompt_cache(req_idx, cache_len, full_slots) + if payload is None: + return None + + src_slots = payload["swa_slots"].long().to(self.kv_buffer.device) + dst_slots = self.swa_allocator.alloc(src_slots.numel()).long().to(self.kv_buffer.device) + for layer_idx in range(self.layer_num): + self.swa_pool.write(layer_idx, dst_slots, self.swa_pool.read(layer_idx, src_slots)) + payload["swa_slots"] = dst_slots.detach().cpu() + return payload + + def detach_swa_for_prompt_cache(self, req_idx: int, swa_payload) -> None: + if ( + swa_payload is None + or self.req_to_swa_indexs is None + or self.req_to_swa_full_indexs is None + or len(swa_payload["positions"]) == 0 + ): + return + req_idx = int(req_idx) + positions = swa_payload["positions"].tolist() + full_slots = swa_payload["full_slots"].tolist() + swa_slots = swa_payload["swa_slots"].tolist() + for pos, full, swa in zip(positions, full_slots, swa_slots): + ring_pos = int(pos) % int(self.sliding_window) + if int(self.req_to_swa_indexs[req_idx, ring_pos].item()) == int(swa) and int( + self.req_to_swa_full_indexs[req_idx, ring_pos].item() + ) == int(full): + self.req_to_swa_indexs[req_idx, ring_pos] = self.swa_pool.HOLD_TOKEN_MEMINDEX + self.req_to_swa_full_indexs[req_idx, ring_pos] = -1 + return + + def restore_swa_from_prompt_cache(self, swa_payload) -> None: + if swa_payload is None or len(swa_payload["full_slots"]) == 0: + return + full_slots = swa_payload["full_slots"].long().to(self.kv_buffer.device) + swa_slots = swa_payload["swa_slots"].long().to(self.kv_buffer.device) + self.full_to_swa_indexs[full_slots] = swa_slots.to(torch.int32) + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + return + + def free_swa_prompt_cache(self, swa_payload) -> None: + if swa_payload is None or len(swa_payload["swa_slots"]) == 0: + return + swa_slots = torch.unique(swa_payload["swa_slots"].long()).detach().cpu() + self.swa_allocator.free(swa_slots) + full_slots = swa_payload["full_slots"].long().to(self.kv_buffer.device) + mapped = self.full_to_swa_indexs[full_slots].long() + expected = swa_payload["swa_slots"].long().to(self.kv_buffer.device) + same = mapped == expected + if same.any(): + self.full_to_swa_indexs[full_slots[same]] = -1 + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + return + def _keep_last_swa_writes(self, swa_slots: torch.Tensor, packed: torch.Tensor): """Drop duplicate SWA writes generated by long prefill ring reuse.""" if swa_slots.numel() <= 1: @@ -634,7 +829,11 @@ def pack_mla_kv_to_cache( if req_idx is None or positions is None: swa_slots = self._identity_swa_slots(mem_index).to(kv.device) else: - swa_slots = self.ensure_swa_slots(req_idx, positions, mem_index).to(kv.device) + pending_slots = self._get_pending_prefill_swa_slots(req_idx, positions, mem_index) + if pending_slots is None: + swa_slots = self.ensure_swa_slots(req_idx, positions, mem_index).to(kv.device) + else: + swa_slots = pending_slots.to(kv.device) swa_slots, packed = self._keep_last_swa_writes(swa_slots, packed) if swa_slots.numel() == 0: return @@ -712,6 +911,7 @@ def free_all(self): if getattr(self, "req_to_swa_indexs", None) is not None: self.req_to_swa_indexs.fill_(self.swa_pool.HOLD_TOKEN_MEMINDEX) self.req_to_swa_full_indexs.fill_(-1) + self._pending_prefill_swa = {} if self.c4_pool is not None: self.c4_pool.free_all() if self.c128_pool is not None: diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 1cdea03381..606469d48e 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -1,5 +1,6 @@ import torch import collections +from dataclasses import dataclass from lightllm.common.linear_att_cache_manager.config_objs import LinearAttCacheConfig from lightllm.utils.log_utils import init_logger @@ -23,6 +24,33 @@ logger = init_logger(__name__) +@dataclass +class DeepseekV4PromptCachePayload: + cache_len: int + c4_slots: Optional[torch.Tensor] = None + c128_slots: Optional[torch.Tensor] = None + c4_state: Optional[torch.Tensor] = None + c4_state_pool: Optional[torch.Tensor] = None + c4_indexer_state: Optional[torch.Tensor] = None + c4_indexer_state_pool: Optional[torch.Tensor] = None + swa: Optional[dict] = None + + +class DeepseekV4PromptCacheValueOps: + def __init__(self, req_manager: "DeepseekV4ReqManager"): + self.req_manager = req_manager + + def slice(self, payload: DeepseekV4PromptCachePayload, start: int, end: int): + return self.req_manager.slice_prompt_cache_payload(payload, start, end) + + def concat(self, payloads: List[DeepseekV4PromptCachePayload]): + return self.req_manager.concat_prompt_cache_payloads(payloads) + + def free(self, payload: DeepseekV4PromptCachePayload): + self.req_manager.free_prompt_cache_payload(payload) + return + + class _ReqNode: def __init__(self, index): self.index = index @@ -594,6 +622,233 @@ def get_c4_indexer_state_pool_for_req(self, layer_index: int, req_idx: int) -> t local = self.layer_to_c4_idx[layer_index] return self.req_to_c4_indexer_state_pool.buffer[local, req_idx] + def get_prompt_cache_value_ops(self): + return DeepseekV4PromptCacheValueOps(self) + + def get_prompt_cache_page_size(self): + return 128 + + def _slice_cpu_slots(self, slots: Optional[torch.Tensor], start: int, end: int, ratio: int): + if slots is None: + return None + return slots[start // ratio : end // ratio].clone() + + def _slice_swa_payload(self, swa_payload, start: int, end: int): + if swa_payload is None: + return None + positions = swa_payload["positions"] + mask = (positions >= start) & (positions < end) + if not bool(mask.any()): + return None + return { + "positions": positions[mask].clone(), + "full_slots": swa_payload["full_slots"][mask].clone(), + "swa_slots": swa_payload["swa_slots"][mask].clone(), + } + + def slice_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload, start: int, end: int): + start = int(start) + end = int(end) + # c4/c128/indexer-K slots are true historical KV and can be sliced by ratio. + # compressor running state only describes the payload end boundary; it is valid + # for a slice only when that slice keeps the original end boundary. + keep_end_state = end == payload.cache_len + return DeepseekV4PromptCachePayload( + cache_len=end - start, + c4_slots=self._slice_cpu_slots(payload.c4_slots, start, end, 4), + c128_slots=self._slice_cpu_slots(payload.c128_slots, start, end, 128), + c4_state=payload.c4_state.clone() if keep_end_state and payload.c4_state is not None else None, + c4_state_pool=payload.c4_state_pool.clone() + if keep_end_state and payload.c4_state_pool is not None + else None, + c4_indexer_state=payload.c4_indexer_state.clone() + if keep_end_state and payload.c4_indexer_state is not None + else None, + c4_indexer_state_pool=payload.c4_indexer_state_pool.clone() + if keep_end_state and payload.c4_indexer_state_pool is not None + else None, + swa=self._slice_swa_payload(payload.swa, start, end), + ) + + def concat_prompt_cache_payloads(self, payloads: List[DeepseekV4PromptCachePayload]): + if len(payloads) == 0: + return None + c4_slots = [p.c4_slots for p in payloads if p.c4_slots is not None and len(p.c4_slots) > 0] + c128_slots = [p.c128_slots for p in payloads if p.c128_slots is not None and len(p.c128_slots) > 0] + last = payloads[-1] + return DeepseekV4PromptCachePayload( + cache_len=sum(p.cache_len for p in payloads), + c4_slots=torch.cat(c4_slots, dim=0) if c4_slots else None, + c128_slots=torch.cat(c128_slots, dim=0) if c128_slots else None, + c4_state=last.c4_state, + c4_state_pool=last.c4_state_pool, + c4_indexer_state=last.c4_indexer_state, + c4_indexer_state_pool=last.c4_indexer_state_pool, + swa=last.swa, + ) + + def build_prompt_cache_payload( + self, + req_idx: int, + cache_len: int, + clone_swa: bool = False, + ) -> DeepseekV4PromptCachePayload: + assert self.mem_manager is not None + cache_len = int(cache_len) + full_slots = self.req_to_token_indexs[req_idx, :cache_len].detach().cpu() + c4_count = cache_len // 4 + c128_count = cache_len // 128 + c4_slots = self.req_to_c4_indexs[req_idx, :c4_count].detach().cpu().clone() if c4_count > 0 else None + c128_slots = self.req_to_c128_indexs[req_idx, :c128_count].detach().cpu().clone() if c128_count > 0 else None + if clone_swa: + swa_payload = self.mem_manager.clone_swa_for_prompt_cache(req_idx, cache_len, full_slots) + else: + swa_payload = self.mem_manager.snapshot_swa_for_prompt_cache(req_idx, cache_len, full_slots) + return DeepseekV4PromptCachePayload( + cache_len=cache_len, + c4_slots=c4_slots, + c128_slots=c128_slots, + c4_state=self.req_to_c4_state.buffer[:, req_idx].detach().clone() if self.n_c4 > 0 else None, + c4_state_pool=self.req_to_c4_state_pool.buffer[:, req_idx].detach().clone() if self.n_c4 > 0 else None, + c4_indexer_state=self.req_to_c4_indexer_state.buffer[:, req_idx].detach().clone() + if self.n_c4 > 0 + else None, + c4_indexer_state_pool=self.req_to_c4_indexer_state_pool.buffer[:, req_idx].detach().clone() + if self.n_c4 > 0 + else None, + swa=swa_payload, + ) + + def detach_prompt_cache_payload_from_req(self, req_idx: int, payload: DeepseekV4PromptCachePayload): + if payload is not None and self.mem_manager is not None: + self.mem_manager.detach_swa_for_prompt_cache(req_idx, payload.swa) + return + + def free_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload): + if payload is None or self.mem_manager is None: + return + if payload.c4_slots is not None and len(payload.c4_slots) > 0: + self.mem_manager.free_c4(payload.c4_slots) + if payload.c128_slots is not None and len(payload.c128_slots) > 0: + self.mem_manager.free_c128(payload.c128_slots) + self.mem_manager.free_swa_prompt_cache(payload.swa) + return + + def release_prompt_cache_detached_swa( + self, + payload: DeepseekV4PromptCachePayload, + keep_payload: Optional[DeepseekV4PromptCachePayload] = None, + ): + if payload is None or payload.swa is None or self.mem_manager is None: + return + old_swa = payload.swa + if keep_payload is None or keep_payload.swa is None: + self.mem_manager.free_swa_prompt_cache(old_swa) + return + + old_slots = old_swa["swa_slots"].long() + keep_slots = keep_payload.swa["swa_slots"].long() + if old_slots.numel() == 0: + return + if keep_slots.numel() == 0: + self.mem_manager.free_swa_prompt_cache(old_swa) + return + + release_mask = ~torch.isin(old_slots, keep_slots) + if not release_mask.any(): + return + release_payload = { + "full_slots": old_swa["full_slots"][release_mask].clone(), + "swa_slots": old_swa["swa_slots"][release_mask].clone(), + } + self.mem_manager.free_swa_prompt_cache(release_payload) + return + + def _reset_c128_for_prompt_cache(self, req_idx: int): + if self.n_c128 > 0: + self._reset_compress_cache_req(self.req_to_c128_state, req_idx) + self._reset_state_pool_req(self.req_to_c128_state_pool, req_idx) + return + + def rebuild_runtime_state_for_req(self, req_idx: int): + state_map = self._runtime_states[req_idx] + state_map.clear() + for layer_index, ratio in enumerate(self.compress_rates): + if ratio == 4: + cstate_kv, cstate_score = self.get_compress_state_for_req(layer_index, req_idx) + idx_state = self.get_c4_indexer_compress_state(layer_index) + state_map[layer_index] = { + "cstate_kv": cstate_kv, + "cstate_score": cstate_score, + "idx_cstate_kv": idx_state[req_idx, 0], + "idx_cstate_score": idx_state[req_idx, 1], + } + elif ratio == 128: + cstate_kv, cstate_score = self.get_compress_state_for_req(layer_index, req_idx) + state_map[layer_index] = { + "cstate_kv": cstate_kv, + "cstate_score": cstate_score, + } + return + + def restore_prompt_cache_payload(self, req_idx: int, payload: DeepseekV4PromptCachePayload): + assert self.mem_manager is not None + cache_len = int(payload.cache_len) + c4_count = cache_len // 4 + c128_count = cache_len // 128 + if c4_count > 0: + assert payload.c4_slots is not None and len(payload.c4_slots) == c4_count + self.req_to_c4_indexs[req_idx, :c4_count] = payload.c4_slots.cuda(non_blocking=True) + if c128_count > 0: + assert payload.c128_slots is not None and len(payload.c128_slots) == c128_count + self.req_to_c128_indexs[req_idx, :c128_count] = payload.c128_slots.cuda(non_blocking=True) + self._c4_entry_counts[req_idx] = c4_count + self._c128_entry_counts[req_idx] = c128_count + + if self.n_c4 > 0: + if payload.c4_state is None or payload.c4_indexer_state is None: + raise RuntimeError("DeepSeek-V4 prompt cache hit is missing c4 running state") + self.req_to_c4_state.buffer[:, req_idx].copy_(payload.c4_state) + self.req_to_c4_indexer_state.buffer[:, req_idx].copy_(payload.c4_indexer_state) + if payload.c4_state_pool is not None: + self.req_to_c4_state_pool.buffer[:, req_idx].copy_(payload.c4_state_pool) + if payload.c4_indexer_state_pool is not None: + self.req_to_c4_indexer_state_pool.buffer[:, req_idx].copy_(payload.c4_indexer_state_pool) + self._reset_c128_for_prompt_cache(req_idx) + self.mem_manager.restore_swa_from_prompt_cache(payload.swa) + self.rebuild_runtime_state_for_req(req_idx) + return + + def pop_prompt_cache_free_compress_indices( + self, + req_idx: int, + keep_len: int, + duplicate_start_len: Optional[int] = None, + duplicate_end_len: Optional[int] = None, + ): + def collect(table, cur_count, ratio): + ranges = [] + if duplicate_start_len is not None and duplicate_end_len is not None: + dup_start = duplicate_start_len // ratio + dup_end = duplicate_end_len // ratio + if dup_end > dup_start: + ranges.append((dup_start, dup_end)) + keep_count = keep_len // ratio + if cur_count > keep_count: + ranges.append((keep_count, cur_count)) + parts = [table[req_idx, s:e].clone() for s, e in ranges if e > s] + return torch.cat(parts, dim=0) if parts else None + + c4 = collect(self.req_to_c4_indexs, self._c4_entry_counts[req_idx], 4) + c128 = collect(self.req_to_c128_indexs, self._c128_entry_counts[req_idx], 128) + if self._c4_entry_counts[req_idx] > 0: + self.req_to_c4_indexs[req_idx, : self._c4_entry_counts[req_idx]].fill_(0) + if self._c128_entry_counts[req_idx] > 0: + self.req_to_c128_indexs[req_idx, : self._c128_entry_counts[req_idx]].fill_(0) + self._c4_entry_counts[req_idx] = 0 + self._c128_entry_counts[req_idx] = 0 + return c4, c128 + def free( self, free_req_indexes, diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 98dee7fd8a..81b45299f9 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -593,19 +593,25 @@ def _indexer_topk(self, idx_q, idx_comp, idx_weight, positions_1based, offset, i if k == 0: return torch.empty((idx_q.shape[0], 0), device=idx_q.device, dtype=torch.long) - scores = torch.einsum("thd,nd->thn", idx_q.float(), idx_comp.float()) - scores = F.relu(scores) * self.indexer_score_scale - index_scores = (scores * idx_weight.unsqueeze(-1)).sum(dim=1) - if self.tp_world_size_ > 1: - all_reduce( - index_scores, - op=dist.ReduceOp.SUM, - group=infer_state.dist_group, - async_op=False, - ) - - causal_threshold = positions_1based // 4 - top = self._indexer_topk_kernel(index_scores, causal_threshold, k) + top_chunks = [] + heads = max(1, idx_q.shape[1]) + max_score_elems = 16 * 1024 * 1024 + chunk_size = max(1, min(idx_q.shape[0], max_score_elems // max(1, heads * ncomp))) + for start in range(0, idx_q.shape[0], chunk_size): + end = min(idx_q.shape[0], start + chunk_size) + scores = torch.einsum("thd,nd->thn", idx_q[start:end].float(), idx_comp.float()) + scores = F.relu(scores) * self.indexer_score_scale + index_scores = (scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1) + if self.tp_world_size_ > 1: + all_reduce( + index_scores, + op=dist.ReduceOp.SUM, + group=infer_state.dist_group, + async_op=False, + ) + causal_threshold = positions_1based[start:end] // 4 + top_chunks.append(self._indexer_topk_kernel(index_scores, causal_threshold, k)) + top = torch.cat(top_chunks, dim=0) valid = top >= 0 return torch.where(valid, top + offset, torch.full_like(top, -1)) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 5a92a339cb..4c23001bce 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -133,10 +133,6 @@ def normal_or_p_d_start(args): raise NotImplementedError("DeepSeek-V4 EP MoE is not supported yet; use TP for now.") if "prompt_cache_kv_buffer" in get_config_json(args.model_dir): raise NotImplementedError("DeepSeek-V4 prompt_cache_kv_buffer is not supported yet.") - if not args.disable_dynamic_prompt_cache: - logger.info("DeepSeek-V4 runtime state does not support radix prompt cache yet; disabling it.") - args.disable_dynamic_prompt_cache = True - args.use_dynamic_prompt_cache = False if args.enable_cpu_cache: # 生成一个用于创建cpu kv cache的共享内存id。 diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index 21e26c5854..8be5198eb3 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -2,7 +2,7 @@ import torch import numpy as np import collections -from typing import Tuple, Dict, Set, List, Optional, Union +from typing import Any, Tuple, Dict, Set, List, Optional, Union from sortedcontainers import SortedSet from .shared_arr import SharedArray @@ -25,6 +25,7 @@ def __init__(self): self.parent: TreeNode = None self.token_id_key: torch.Tensor = None self.token_mem_index_value: torch.Tensor = None # 用于记录存储的 token_index 为每个元素在 token mem 中的index位置 + self.token_extra_value: Any = None self.ref_counter = 0 self.time_id = time_gen.generate_time_id() # 用于标识时间周期 @@ -34,14 +35,17 @@ def __init__(self): def get_compare_key(self): return (0 if self.ref_counter == 0 else 1, len(self.children), self.time_id) - def split_node(self, prefix_len): + def split_node(self, prefix_len, child_key_fn=None, extra_value_ops=None): split_parent_node = TreeNode() split_parent_node.parent = self.parent - split_parent_node.parent.children[self.token_id_key[0].item()] = split_parent_node + split_parent_node.parent.children[child_key_fn(self.token_id_key)] = split_parent_node split_parent_node.token_id_key = self.token_id_key[0:prefix_len] split_parent_node.token_mem_index_value = self.token_mem_index_value[0:prefix_len] + if self.token_extra_value is not None and extra_value_ops is not None: + split_parent_node.token_extra_value = extra_value_ops.slice(self.token_extra_value, 0, prefix_len) + self.token_extra_value = extra_value_ops.slice(self.token_extra_value, prefix_len, len(self.token_id_key)) split_parent_node.children = {} - split_parent_node.children[self.token_id_key[prefix_len].item()] = self + split_parent_node.children[child_key_fn(self.token_id_key[prefix_len:])] = self split_parent_node.ref_counter = self.ref_counter new_len = len(split_parent_node.token_mem_index_value) @@ -56,11 +60,12 @@ def split_node(self, prefix_len): self.node_prefix_total_len = self.parent.node_prefix_total_len + new_len return split_parent_node - def add_and_return_new_child(self, token_id_key, token_mem_index_value): + def add_and_return_new_child(self, token_id_key, token_mem_index_value, token_extra_value=None, child_key=None): child = TreeNode() child.token_id_key = token_id_key child.token_mem_index_value = token_mem_index_value - first_token_key = child.token_id_key[0].item() + child.token_extra_value = token_extra_value + first_token_key = child.token_id_key[0].item() if child_key is None else child_key assert first_token_key not in self.children.keys() self.children[first_token_key] = child child.parent = self @@ -71,9 +76,17 @@ def add_and_return_new_child(self, token_id_key, token_mem_index_value): return child def remove_child(self, child_node: "TreeNode"): - del self.children[child_node.token_id_key[0].item()] - child_node.parent = None - return + child_key = child_node.token_id_key[0].item() + if child_key in self.children: + del self.children[child_key] + child_node.parent = None + return + for key, value in list(self.children.items()): + if value is child_node: + del self.children[key] + child_node.parent = None + return + raise KeyError("child node not found") def update_time(self): self.time_id = time_gen.generate_time_id() @@ -103,12 +116,22 @@ class RadixCache: unique_name 主要用于解决单机,多实列部署时的shm冲突 """ - def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None): + def __init__( + self, + unique_name, + total_token_num, + rank_in_node, + mem_manager=None, + page_size: int = 1, + extra_value_ops=None, + ): from lightllm.common.kv_cache_mem_manager import MemoryManager self.mem_manager: MemoryManager = mem_manager self._key_dtype = torch.int64 self._value_dtype = torch.int64 + self.page_size = max(1, int(page_size)) + self.extra_value_ops = extra_value_ops self.root_node = TreeNode() self.root_node.token_id_key = torch.zeros((0,), device="cpu", dtype=self._key_dtype) @@ -125,30 +148,66 @@ def __init__(self, unique_name, total_token_num, rank_in_node, mem_manager=None) ) self.tree_total_tokens_num.arr[0] = 0 - def insert(self, key, value=None) -> Tuple[int, Optional[TreeNode]]: + def _align_len(self, length: int) -> int: + if self.page_size <= 1: + return int(length) + return int(length) // self.page_size * self.page_size + + def align_len(self, length: int) -> int: + return self._align_len(length) + + def _child_key(self, key: torch.Tensor): + if self.page_size <= 1: + return key[0].item() + return tuple(key[: self.page_size].tolist()) + + def _match_len(self, key: torch.Tensor, node_key: torch.Tensor) -> int: + prefix_len = match(key, node_key) + return self._align_len(prefix_len) + + def _slice_extra(self, extra_value, start: int, end: int): + if extra_value is None: + return None + assert self.extra_value_ops is not None + return self.extra_value_ops.slice(extra_value, start, end) + + def _concat_extra(self, values: list): + values = [v for v in values if v is not None] + if len(values) == 0: + return None + assert self.extra_value_ops is not None + return self.extra_value_ops.concat(values) + + def insert(self, key, value=None, extra_value=None) -> Tuple[int, Optional[TreeNode]]: if value is None: value = key + align_len = self._align_len(len(key)) + key = key[:align_len] + value = value[:align_len] + if extra_value is not None: + extra_value = self._slice_extra(extra_value, 0, align_len) + assert len(key) == len(value) # and len(key) >= 1 if len(key) == 0: return 0, None - return self._insert_helper(self.root_node, key, value) + return self._insert_helper(self.root_node, key, value, extra_value) - def _insert_helper(self, node: TreeNode, key, value) -> Tuple[int, Optional[TreeNode]]: + def _insert_helper(self, node: TreeNode, key, value, extra_value) -> Tuple[int, Optional[TreeNode]]: handle_stack = collections.deque() update_list = collections.deque() - handle_stack.append((node, key, value)) + handle_stack.append((node, key, value, extra_value)) ans_prefix_len = 0 ans_node = None while len(handle_stack) != 0: - node, key, value = handle_stack.popleft() - ans_tuple = self._insert_helper_no_recursion(node=node, key=key, value=value) - if len(ans_tuple) == 4: - (_prefix_len, new_node, new_key, new_value) = ans_tuple + node, key, value, extra_value = handle_stack.popleft() + ans_tuple = self._insert_helper_no_recursion(node=node, key=key, value=value, extra_value=extra_value) + if len(ans_tuple) == 5: + (_prefix_len, new_node, new_key, new_value, new_extra_value) = ans_tuple ans_prefix_len += _prefix_len - handle_stack.append((new_node, new_key, new_value)) + handle_stack.append((new_node, new_key, new_value, new_extra_value)) else: _prefix_len, ans_node = ans_tuple ans_prefix_len += _prefix_len @@ -166,15 +225,15 @@ def _insert_helper(self, node: TreeNode, key, value) -> Tuple[int, Optional[Tree return ans_prefix_len, ans_node def _insert_helper_no_recursion( - self, node: TreeNode, key: torch.Tensor, value: torch.Tensor - ) -> Union[Tuple[int, Optional[TreeNode]], Tuple[int, TreeNode, torch.Tensor, torch.Tensor]]: + self, node: TreeNode, key: torch.Tensor, value: torch.Tensor, extra_value=None + ) -> Union[Tuple[int, Optional[TreeNode]], Tuple[int, TreeNode, torch.Tensor, torch.Tensor, Any]]: if node.is_leaf(): self.evict_tree_set.discard(node) - first_key_id = key[0].item() + first_key_id = self._child_key(key) if first_key_id in node.children.keys(): child: TreeNode = node.children[first_key_id] - prefix_len = match(key, child.token_id_key) + prefix_len = self._match_len(key, child.token_id_key) if prefix_len == len(key): if prefix_len == len(child.token_id_key): if child.is_leaf(): @@ -184,10 +243,14 @@ def _insert_helper_no_recursion( self.evict_tree_set.add(child) return prefix_len, child elif prefix_len < len(child.token_id_key): + if prefix_len == 0: + return 0, node if child.is_leaf(): self.evict_tree_set.discard(child) - split_parent_node = child.split_node(prefix_len) + split_parent_node = child.split_node( + prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops + ) if split_parent_node.is_leaf(): self.evict_tree_set.add(split_parent_node) @@ -199,13 +262,23 @@ def _insert_helper_no_recursion( assert False, "can not run to here" elif prefix_len < len(key) and prefix_len < len(child.token_id_key): + if prefix_len == 0: + return 0, node if child.is_leaf(): self.evict_tree_set.discard(child) + new_extra_value = self._slice_extra(extra_value, prefix_len, len(key)) key = key[prefix_len:] value = value[prefix_len:] - split_parent_node = child.split_node(prefix_len) - new_node = split_parent_node.add_and_return_new_child(key, value) + split_parent_node = child.split_node( + prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops + ) + new_node = split_parent_node.add_and_return_new_child( + key, + value, + token_extra_value=new_extra_value, + child_key=self._child_key(key), + ) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) @@ -218,12 +291,23 @@ def _insert_helper_no_recursion( self.evict_tree_set.add(child) return prefix_len, new_node elif prefix_len < len(key) and prefix_len == len(child.token_id_key): - return (prefix_len, child, key[prefix_len:], value[prefix_len:]) + return ( + prefix_len, + child, + key[prefix_len:], + value[prefix_len:], + self._slice_extra(extra_value, prefix_len, len(key)), + ) else: assert False, "can not run to here" else: - new_node = node.add_and_return_new_child(key, value) + new_node = node.add_and_return_new_child( + key, + value, + token_extra_value=extra_value, + child_key=first_key_id, + ) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) if new_node.is_leaf(): @@ -231,7 +315,9 @@ def _insert_helper_no_recursion( return 0, new_node def match_prefix(self, key, update_refs=False): - assert len(key) != 0 + key = key[: self._align_len(len(key))] + if len(key) == 0: + return None, 0, None ans_value_list = [] tree_node = self._match_prefix_helper(self.root_node, key, ans_value_list, update_refs=update_refs) if tree_node != self.root_node: @@ -290,20 +376,24 @@ def _match_prefix_helper_no_recursion( if len(key) == 0: return node - first_key_id = key[0].item() + first_key_id = self._child_key(key) if first_key_id not in node.children.keys(): return node else: child = node.children[first_key_id] - prefix_len = match(key, child.token_id_key) + prefix_len = self._match_len(key, child.token_id_key) if prefix_len == len(child.token_id_key): ans_value_list.append(child.token_mem_index_value) return (child, key[prefix_len:]) elif prefix_len < len(child.token_id_key): + if prefix_len == 0: + return node if child.is_leaf(): self.evict_tree_set.discard(child) - split_parent_node = child.split_node(prefix_len) + split_parent_node = child.split_node( + prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops + ) ans_value_list.append(split_parent_node.token_mem_index_value) if update_refs: @@ -334,6 +424,8 @@ def evict(self, need_remove_tokens, evict_callback): ), "error evict tree node state" num_evicted += len(node.token_mem_index_value) evict_callback(node.token_mem_index_value) + if self.extra_value_ops is not None and node.token_extra_value is not None: + self.extra_value_ops.free(node.token_extra_value) # update total token num self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value) parent_node: TreeNode = node.parent @@ -369,11 +461,12 @@ def _try_merge(self, child_node: TreeNode) -> Optional[TreeNode]: child_node.token_mem_index_value = torch.cat( [parent_node.token_mem_index_value, child_node.token_mem_index_value] ) + child_node.token_extra_value = self._concat_extra([parent_node.token_extra_value, child_node.token_extra_value]) child_node.node_value_len = len(child_node.token_mem_index_value) child_node.time_id = max(parent_node.time_id, child_node.time_id) grandparent_node = parent_node.parent - key_in_grandparent = parent_node.token_id_key[0].item() + key_in_grandparent = self._child_key(parent_node.token_id_key) grandparent_node.children[key_in_grandparent] = child_node child_node.parent = grandparent_node @@ -469,6 +562,19 @@ def get_mem_index_value_by_node(self, node: TreeNode) -> Optional[torch.Tensor]: ans_list.reverse() return torch.concat(ans_list, dim=0) + def get_extra_value_by_node(self, node: TreeNode): + if node is None or self.extra_value_ops is None: + return None + + ans_list = [] + while node is not None: + if node.token_extra_value is not None: + ans_list.append(node.token_extra_value) + node = node.parent + + ans_list.reverse() + return self._concat_extra(ans_list) + def get_refed_tokens_num(self): return self.refed_tokens_num.arr[0] diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index abeb8d61e9..3fd6e0463a 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -131,7 +131,8 @@ def free_a_req_mem( free_c4_index: Optional[List] = None, free_c128_index: Optional[List] = None, ): - if hasattr(self.req_manager, "pop_compress_indices_for_req"): + is_dsv4_req_manager = hasattr(self.req_manager, "build_prompt_cache_payload") + if hasattr(self.req_manager, "pop_compress_indices_for_req") and not is_dsv4_req_manager: c4, c128 = self.req_manager.pop_compress_indices_for_req(req.req_idx) if c4 is not None and free_c4_index is not None: free_c4_index.append(c4) @@ -141,9 +142,24 @@ def free_a_req_mem( if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) + if is_dsv4_req_manager: + c4, c128 = self.req_manager.pop_compress_indices_for_req(req.req_idx) + if c4 is not None and free_c4_index is not None: + free_c4_index.append(c4) + if c128 is not None and free_c128_index is not None: + free_c128_index.append(c128) + self.req_manager.clear_runtime_state(req.req_idx) else: if not self.is_linear_att_mixed_model: - self._full_att_free_req(free_token_index=free_token_index, req=req) + if is_dsv4_req_manager: + self._dsv4_full_att_free_req( + free_token_index=free_token_index, + req=req, + free_c4_index=free_c4_index, + free_c128_index=free_c128_index, + ) + else: + self._full_att_free_req(free_token_index=free_token_index, req=req) else: self._linear_att_free_req(free_token_index=free_token_index, req=req) assert len(req.linear_att_len_to_big_page_id) == 0 @@ -151,6 +167,11 @@ def free_a_req_mem( req.shm_req.shm_cur_kv_len = req.cur_kv_len return + def _append_free_token_index(self, free_token_index: List, tensor: torch.Tensor): + if tensor.numel() > 0: + free_token_index.append(tensor) + return + def _full_att_free_req(self, free_token_index: List, req: "InferReq"): input_token_ids = req.get_input_token_ids() key = torch.tensor(input_token_ids[0 : req.cur_kv_len], dtype=torch.int64, device="cpu") @@ -166,6 +187,86 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): req.shared_kv_node = None return + def _dsv4_full_att_free_req( + self, + free_token_index: List, + req: "InferReq", + free_c4_index: Optional[List] = None, + free_c128_index: Optional[List] = None, + ): + if req.cur_kv_len == 0: + free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0:0]) + return + + old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len + cache_len = self.radix_cache.align_len(req.cur_kv_len) + inserted_len = old_prefix_len + duplicate_prefix_len = old_prefix_len + inserted_payload = None + pending_payload = getattr(req, "prompt_cache_snapshot_payload", None) + pending_cache_len = getattr(req, "prompt_cache_snapshot_len", 0) + + # The current V4 runtime state is only guaranteed to describe the current + # sequence end. Cache aligned current ends; leave unaligned tails uncached. + if pending_payload is not None and pending_cache_len > old_prefix_len: + cache_len = pending_cache_len + input_token_ids = req.get_input_token_ids() + key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") + value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() + duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value, extra_value=pending_payload) + inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len + if inserted_len == cache_len: + inserted_payload = pending_payload + else: + self.req_manager.release_prompt_cache_detached_swa(pending_payload) + pending_payload = None + inserted_len = old_prefix_len + duplicate_prefix_len = old_prefix_len + elif cache_len == req.cur_kv_len and cache_len > old_prefix_len: + input_token_ids = req.get_input_token_ids() + key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") + value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() + payload = self.req_manager.build_prompt_cache_payload(req.req_idx, cache_len) + duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value, extra_value=payload) + inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len + if inserted_len == cache_len: + inserted_payload = payload + self.req_manager.detach_prompt_cache_payload_from_req(req.req_idx, inserted_payload) + else: + inserted_len = old_prefix_len + duplicate_prefix_len = old_prefix_len + + if ( + pending_payload is not None + and inserted_payload is not pending_payload + and pending_cache_len <= old_prefix_len + ): + self.req_manager.release_prompt_cache_detached_swa(pending_payload) + req.prompt_cache_snapshot_payload = None + req.prompt_cache_snapshot_len = 0 + dense_row = self.req_manager.req_to_token_indexs[req.req_idx] + self._append_free_token_index(free_token_index, dense_row[old_prefix_len:duplicate_prefix_len]) + self._append_free_token_index(free_token_index, dense_row[inserted_len : req.cur_kv_len]) + if len(free_token_index) == 0: + free_token_index.append(dense_row[0:0]) + + c4, c128 = self.req_manager.pop_prompt_cache_free_compress_indices( + req.req_idx, + keep_len=inserted_len, + duplicate_start_len=old_prefix_len, + duplicate_end_len=duplicate_prefix_len, + ) + if c4 is not None and free_c4_index is not None: + free_c4_index.append(c4) + if c128 is not None and free_c128_index is not None: + free_c128_index.append(c128) + + if req.shared_kv_node is not None: + assert req.shared_kv_node.node_prefix_total_len <= max(inserted_len, old_prefix_len) + self.radix_cache.dec_node_ref_counter(req.shared_kv_node) + req.shared_kv_node = None + return + def _linear_att_free_req(self, free_token_index: List, req: "InferReq"): assert g_infer_context.is_linear_att_mixed_model is True args = get_env_start_args() @@ -637,6 +738,8 @@ def _init_all_state(self): g_infer_context.req_manager.req_sampling_params_manager.init_req_sampling_params(self) if hasattr(g_infer_context.req_manager, "init_compress_state"): g_infer_context.req_manager.init_compress_state(req_idx=self.req_idx) + self.prompt_cache_snapshot_len = 0 + self.prompt_cache_snapshot_payload = None self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() # token healing mode 才被使用的管理对象 @@ -672,6 +775,11 @@ def _match_radix_cache(self): ready_cache_len = share_node.node_prefix_total_len # 从 cpu 到 gpu 是流内阻塞操作 g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor + if hasattr(g_infer_context.req_manager, "restore_prompt_cache_payload"): + payload = g_infer_context.radix_cache.get_extra_value_by_node(share_node) + if payload is None: + raise RuntimeError("DeepSeek-V4 radix cache hit is missing prompt-cache payload") + g_infer_context.req_manager.restore_prompt_cache_payload(self.req_idx, payload) self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 @@ -843,6 +951,7 @@ def get_input_token_ids(self): def get_chuncked_input_token_ids(self): chunked_start = self.cur_kv_len chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) + chunked_end = self._align_chuncked_end_for_prompt_cache(chunked_start, chunked_end) return self.shm_req.shm_prompt_ids.arr[0:chunked_end] def get_chuncked_input_token_ids_for_linear_att(self): @@ -863,6 +972,17 @@ def get_chuncked_input_token_ids_for_linear_att(self): def get_chuncked_input_token_len(self): chunked_start = self.cur_kv_len chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) + return self._align_chuncked_end_for_prompt_cache(chunked_start, chunked_end) + + def _align_chuncked_end_for_prompt_cache(self, chunked_start: int, chunked_end: int): + radix_cache = g_infer_context.radix_cache + page_size = getattr(radix_cache, "page_size", 1) if radix_cache is not None else 1 + if page_size <= 1 or self.sampling_param.disable_prompt_cache: + return chunked_end + prompt_end = self.shm_req.input_len + next_page_end = ((int(chunked_start) // page_size) + 1) * page_size + if int(chunked_start) < next_page_end < int(chunked_end) and next_page_end <= prompt_end: + return next_page_end return chunked_end def get_chuncked_input_token_len_for_linear_att(self): 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 0220dc87fb..74fdb1e87b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -171,6 +171,8 @@ def init_model(self, kvargs): self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) + if hasattr(self.model.req_manager, "build_prompt_cache_payload"): + self.support_overlap = False if self.is_linear_att_mixed_model: self.linear_att_cache_manager = LinearAttCacheManager( @@ -182,6 +184,7 @@ def init_model(self, kvargs): if not self.use_dynamic_prompt_cache: self.radix_cache = None + setattr(self.args, "dynamic_prompt_cache_page_size", 1) else: if self.is_linear_att_mixed_model: self.radix_cache = LinearAttPagedRadixCache( @@ -193,12 +196,21 @@ def init_model(self, kvargs): kv_cache_mem_manager=self.model.mem_manager, linear_att_small_page_buffers=self.linear_att_cache_manager, ) + setattr(self.args, "dynamic_prompt_cache_page_size", 1) else: + radix_page_size = 1 + radix_extra_value_ops = None + if hasattr(self.model.req_manager, "get_prompt_cache_value_ops"): + radix_page_size = self.model.req_manager.get_prompt_cache_page_size() + radix_extra_value_ops = self.model.req_manager.get_prompt_cache_value_ops() + setattr(self.args, "dynamic_prompt_cache_page_size", radix_page_size) self.radix_cache = RadixCache( unique_name=get_unique_server_name(), total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, mem_manager=self.model.mem_manager, + page_size=radix_page_size, + extra_value_ops=radix_extra_value_ops, ) if "prompt_cache_kv_buffer" in model_cfg: @@ -701,6 +713,31 @@ def _pre_handle_finished_reqs(self, finished_reqs: List[InferReq]): """ pass + def _maybe_capture_prompt_cache_payload(self, req_obj: InferReq): + if self.radix_cache is None: + return + req_manager = g_infer_context.req_manager + if not hasattr(req_manager, "build_prompt_cache_payload"): + return + if req_obj.sampling_param.disable_prompt_cache: + return + page_size = getattr(self.args, "dynamic_prompt_cache_page_size", 1) + cache_len = int(req_obj.cur_kv_len) + if page_size <= 1 or cache_len <= 0 or cache_len % page_size != 0: + return + if cache_len > req_obj.shm_req.input_len: + return + if getattr(req_obj, "prompt_cache_snapshot_len", 0) >= cache_len: + return + + payload = req_manager.build_prompt_cache_payload(req_obj.req_idx, cache_len, clone_swa=True) + old_payload = getattr(req_obj, "prompt_cache_snapshot_payload", None) + if old_payload is not None: + req_manager.release_prompt_cache_detached_swa(old_payload, keep_payload=payload) + req_obj.prompt_cache_snapshot_len = cache_len + req_obj.prompt_cache_snapshot_payload = payload + return + # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] @@ -748,6 +785,7 @@ def _post_handle( ): req_obj: InferReq = req_obj pack: InferReqUpdatePack = pack + self._maybe_capture_prompt_cache_payload(req_obj) pack.handle( next_token_id=next_token_id, next_token_logprob=next_token_logprob, From 61eed870b00927005103a99102ab5dc7ead157df Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 5 Jun 2026 09:08:32 +0000 Subject: [PATCH 004/214] support cudagraph --- lightllm/common/basemodel/basemodel.py | 44 +- .../deepseek4_mem_manager.py | 66 +++ lightllm/common/req_manager.py | 20 + .../deepseek_v4/layer_infer/attention.py | 47 +- .../deepseek_v4/layer_infer/compressor.py | 64 +++ .../layer_infer/transformer_layer_infer.py | 473 +++++++++++++----- lightllm/models/deepseek_v4/model.py | 30 +- 7 files changed, 615 insertions(+), 129 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index d785991808..8e352519c0 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -577,6 +577,12 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) + if hasattr(self.mem_manager, "prepare_decode_swa_slots"): + self.mem_manager.prepare_decode_swa_slots( + model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes + ) + if hasattr(self.req_manager, "prepare_decode_compress_slots"): + self.req_manager.prepare_decode_compress_slots(model_input.b_req_idx, model_input.b_seq_len) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -598,6 +604,12 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) + if hasattr(self.mem_manager, "prepare_decode_swa_slots"): + self.mem_manager.prepare_decode_swa_slots( + model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes + ) + if hasattr(self.req_manager, "prepare_decode_compress_slots"): + self.req_manager.prepare_decode_compress_slots(model_input.b_req_idx, model_input.b_seq_len) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -633,7 +645,13 @@ def prefill_func(input_tensors, infer_state): handle_token_num = infer_state.input_ids.shape[0] - if self.prefill_graph is not None and self.prefill_graph.can_run(handle_token_num=handle_token_num): + can_run_prefill_graph = self.prefill_graph is not None and self.prefill_graph.can_run( + handle_token_num=handle_token_num + ) + if can_run_prefill_graph and hasattr(self, "_can_run_prefill_cudagraph"): + can_run_prefill_graph = self._can_run_prefill_cudagraph(infer_state, handle_token_num) + + if can_run_prefill_graph: finded_handle_token_num = self.prefill_graph.find_closest_graph_handle_token_num( handle_token_num=handle_token_num ) @@ -846,6 +864,20 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode # 一致,需要按照较高 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) + if hasattr(self.mem_manager, "prepare_decode_swa_slots"): + self.mem_manager.prepare_decode_swa_slots( + padded_model_input0.b_req_idx, padded_model_input0.b_seq_len, padded_model_input0.mem_indexes + ) + self.mem_manager.prepare_decode_swa_slots( + padded_model_input1.b_req_idx, padded_model_input1.b_seq_len, padded_model_input1.mem_indexes + ) + if hasattr(self.req_manager, "prepare_decode_compress_slots"): + self.req_manager.prepare_decode_compress_slots( + padded_model_input0.b_req_idx, padded_model_input0.b_seq_len + ) + self.req_manager.prepare_decode_compress_slots( + padded_model_input1.b_req_idx, padded_model_input1.b_seq_len + ) infer_state0 = self._create_inferstate(padded_model_input0, 0) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -887,6 +919,16 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode 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) + if hasattr(self.mem_manager, "prepare_decode_swa_slots"): + self.mem_manager.prepare_decode_swa_slots( + model_input0.b_req_idx, model_input0.b_seq_len, model_input0.mem_indexes + ) + self.mem_manager.prepare_decode_swa_slots( + model_input1.b_req_idx, model_input1.b_seq_len, model_input1.mem_indexes + ) + if hasattr(self.req_manager, "prepare_decode_compress_slots"): + self.req_manager.prepare_decode_compress_slots(model_input0.b_req_idx, model_input0.b_seq_len) + self.req_manager.prepare_decode_compress_slots(model_input1.b_req_idx, model_input1.b_seq_len) infer_state0 = self._create_inferstate(model_input0, 0) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 6735b2deed..dc708e0790 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -561,6 +561,37 @@ def ensure_swa_slots(self, req_idx: int, positions: torch.Tensor, full_slots: to out[i] = swa return out + def prepare_decode_swa_slots( + self, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + mem_index: torch.Tensor, + ) -> None: + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: + return + + reqs = b_req_idx.detach().cpu().tolist() + seqs = b_seq_len.detach().cpu().tolist() + fulls = mem_index.detach().cpu().tolist() + hold = self.swa_pool.HOLD_TOKEN_MEMINDEX + for req_idx, seq_len, full in zip(reqs, seqs, fulls): + req_idx = int(req_idx) + full = int(full) + if req_idx == self.max_request_num or full == self.HOLD_TOKEN_MEMINDEX: + continue + ring_pos = (int(seq_len) - 1) % int(self.sliding_window) + old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) + old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) + if old_swa == hold: + old_swa = int(self.swa_allocator.alloc(1)[0].item()) + if old_full >= 0 and old_full != full: + self.full_to_swa_indexs[old_full] = -1 + self.req_to_swa_indexs[req_idx, ring_pos] = old_swa + self.req_to_swa_full_indexs[req_idx, ring_pos] = full + self.full_to_swa_indexs[full] = old_swa + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = hold + return + def _reserve_prefill_swa_slots( self, req_idx: int, @@ -839,6 +870,41 @@ def pack_mla_kv_to_cache( return self.swa_pool.write(layer_index, swa_slots, packed) + def pack_decode_mla_kv_to_cache( + self, + layer_index: int, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + mem_index: torch.Tensor, + kv: torch.Tensor, + ): + if kv.shape[0] == 0: + return + packed = self._pack_mla_kv(kv) + if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: + swa_slots = self._identity_swa_slots(mem_index).to(kv.device) + else: + req = b_req_idx.long() + ring = ((b_seq_len.long() - 1) % int(self.sliding_window)).long() + swa_slots = self.req_to_swa_indexs[req, ring].long() + + old_full = self.req_to_swa_full_indexs[req, ring].long() + full_slots = mem_index.long() + old_full = torch.where(old_full >= 0, old_full, full_slots) + self.full_to_swa_indexs[old_full] = torch.full( + old_full.shape, + -1, + dtype=self.full_to_swa_indexs.dtype, + device=old_full.device, + ) + + self.req_to_swa_full_indexs[req, ring] = full_slots.to(torch.int32) + self.full_to_swa_indexs[full_slots] = swa_slots.to(torch.int32) + self.swa_pool.write(layer_index, swa_slots.to(kv.device), packed) + + def gather_mla_kv_from_swa_slots(self, layer_index: int, swa_slots: torch.Tensor) -> torch.Tensor: + return self._unpack_mla_kv(self.swa_pool.read(layer_index, swa_slots.to(self.kv_buffer.device))) + def pack_compressed_kv_to_cache(self, layer_index: int, slots: torch.Tensor, comp: torch.Tensor): if comp.shape[0] == 0: return diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 606469d48e..ca027a63c8 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -542,6 +542,26 @@ def ensure_compress_slots(self, layer_index: int, req_idx: int, entry_start: int return self.ensure_c128_slots(req_idx, entry_start, entry_count) raise AssertionError(f"layer {layer_index} is not a compressed attention layer") + def prepare_decode_compress_slots(self, b_req_idx: torch.Tensor, b_seq_len: torch.Tensor) -> None: + req_list = b_req_idx.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + for req_idx, seq_len in zip(req_list, seq_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + seq_len = int(seq_len) + if self.n_c4 > 0: + required_c4 = seq_len // 4 + old_c4 = self._c4_entry_counts[req_idx] + if required_c4 > old_c4: + self.ensure_c4_slots(req_idx, old_c4, required_c4 - old_c4) + if self.n_c128 > 0: + required_c128 = seq_len // 128 + old_c128 = self._c128_entry_counts[req_idx] + if required_c128 > old_c128: + self.ensure_c128_slots(req_idx, old_c128, required_c128 - old_c128) + return + def pop_compress_indices_for_req(self, req_idx: int): c4_count = self._c4_entry_counts[req_idx] if c4_count > 0: diff --git a/lightllm/models/deepseek_v4/layer_infer/attention.py b/lightllm/models/deepseek_v4/layer_infer/attention.py index a24949696f..8a7428f0dd 100644 --- a/lightllm/models/deepseek_v4/layer_infer/attention.py +++ b/lightllm/models/deepseek_v4/layer_infer/attention.py @@ -46,9 +46,13 @@ def _pad_heads_for_flashmla(q, attn_sink): def _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): - q0 = q[0].float() - kv0 = kv[0].float() - indices = topk_idxs[0].long() + return _torch_sparse_attn_flat(q[0], kv[0], attn_sink, topk_idxs[0], scale).unsqueeze(0) + + +def _torch_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale): + q0 = q.float() + kv0 = kv.float() + indices = topk_idxs.long() valid = (indices >= 0) & (indices < kv0.shape[0]) safe_indices = torch.where(valid, indices, torch.zeros_like(indices)) kv_sel = kv0[safe_indices] @@ -60,7 +64,7 @@ def _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): exp_sink = torch.exp(sink - max_scores) denom = exp_scores.sum(dim=-1) + exp_sink out = torch.einsum("mhk,mkd->mhd", exp_scores / denom.unsqueeze(-1), kv_sel) - return out.unsqueeze(0).to(q.dtype) + return out.to(q.dtype) def vllm_sparse_attn(q, kv, attn_sink, topk_idxs, scale): @@ -77,15 +81,40 @@ def vllm_sparse_attn(q, kv, attn_sink, topk_idxs, scale): if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires bf16 q/kv, got {q.dtype}/{kv.dtype}") + return vllm_sparse_attn_flat(q[0], kv[0], attn_sink, topk_idxs[0], scale).unsqueeze(0) + + +def vllm_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale, already_compact=False): + """FlashMLA sparse attention over a flat KV arena. + + q:[m,h,d], kv:[n,d], topk_idxs:[m,K] int. Indices are global offsets into + the flat kv tensor, so callers can concatenate per-request KV candidates and + run one FlashMLA call for the whole batch. When already_compact=True, each + row must place all valid indices before invalid (-1) entries. + """ + m, h, d = q.shape + if d != 512: + raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires head_dim=512, got {d}") + if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: + raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires bf16 q/kv, got {q.dtype}/{kv.dtype}") + if q.shape[0] == 0: + return q.new_empty((0, h, d)) + if DSV4_DEBUG_TORCH_SPARSE_ATTN: - return _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale) + return _torch_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale) from vllm.third_party.flashmla.flash_mla_interface import flash_mla_sparse_fwd - q_pad, sink_pad, real_heads = _pad_heads_for_flashmla(q[0], attn_sink) - indices, topk_lens = _compact_topk_indices(topk_idxs[0].to(torch.int32), kv.shape[1]) + q_pad, sink_pad, real_heads = _pad_heads_for_flashmla(q, attn_sink) + topk_idxs = topk_idxs.to(torch.int32) + if already_compact: + valid = (topk_idxs >= 0) & (topk_idxs < kv.shape[0]) + indices = topk_idxs.contiguous() + topk_lens = valid.sum(dim=-1).to(torch.int32).contiguous() + else: + indices, topk_lens = _compact_topk_indices(topk_idxs, kv.shape[0]) indices = _pad_topk_for_flashmla(indices).unsqueeze(1) - kv_flat = kv[0].unsqueeze(1).contiguous() + kv_flat = kv.unsqueeze(1).contiguous() out, _, _ = flash_mla_sparse_fwd( q=q_pad, kv=kv_flat, @@ -95,7 +124,7 @@ def vllm_sparse_attn(q, kv, attn_sink, topk_idxs, scale): topk_length=topk_lens, out=None, ) - return out[:, :real_heads].unsqueeze(0).to(q.dtype) + return out[:, :real_heads].to(q.dtype) def build_prefill_topk_idxs(seqlen, window, ratio, n_window, device): diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index c91799f9ee..f51f73829c 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -459,3 +459,67 @@ def compressor_decode_step( eps, dtype, ) + + +def compressor_decode_step_batch( + x_new, + wkv_w, + wgate_w, + norm_w, + ape, + ratio, + head_dim, + rope_dim, + cos_table, + sin_table, + eps, + state_all, + b_req_idx, + start_pos, +): + """Graph-safe batch decode compressor step. + + Mutates ``state_all`` for the selected request rows and returns one candidate + entry per batch row plus a boolean mask telling which rows closed a + compression window. + """ + overlap = ratio == 4 + d = head_dim + dtype = x_new.dtype + req = b_req_idx.long() + pos = start_pos.long() + pos_mod = pos % ratio + + xf = x_new.float() + kv = F.linear(xf, wkv_w.float()) + score = F.linear(xf, wgate_w.float()) + ape.float().index_select(0, pos_mod) + + kv_state = state_all[req, 0].clone() + score_state = state_all[req, 1].clone() + row = pos_mod + (ratio if overlap else 0) + batch_ids = torch.arange(x_new.shape[0], device=x_new.device) + kv_state[batch_ids, row] = kv + score_state[batch_ids, row] = score + + should_compress = ((pos + 1) % ratio) == 0 + if overlap: + kv_cat = torch.cat([kv_state[:, :ratio, :d], kv_state[:, ratio:, d:]], dim=1) + score_cat = torch.cat([score_state[:, :ratio, :d], score_state[:, ratio:, d:]], dim=1) + entry = (kv_cat * torch.softmax(score_cat, dim=1)).sum(dim=1) + shifted_kv_state = kv_state.clone() + shifted_score_state = score_state.clone() + shifted_kv_state[:, :ratio] = kv_state[:, ratio:] + shifted_score_state[:, :ratio] = score_state[:, ratio:] + kv_state = torch.where(should_compress.view(-1, 1, 1), shifted_kv_state, kv_state) + score_state = torch.where(should_compress.view(-1, 1, 1), shifted_score_state, score_state) + else: + entry = (kv_state * torch.softmax(score_state, dim=1)).sum(dim=1) + + state_all[req, 0] = kv_state + state_all[req, 1] = score_state + + entry = _rmsnorm(entry.to(dtype), norm_w, eps) + comp_pos = torch.clamp(pos + 1 - ratio, min=0) + entry_rope = apply_rotary_emb(entry[:, -rope_dim:], cos_table[comp_pos], sin_table[comp_pos]) + entry = torch.cat([entry[:, :-rope_dim], entry_rope], dim=1) + return entry, should_compress diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 81b45299f9..864104fd32 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,19 +1,14 @@ -import os - import torch import torch.nn.functional as F import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl from lightllm.distributed.communication_op import all_reduce from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from .hyper_connection import hc_pre, hc_post from ..triton_kernel.rotary_emb import apply_rotary_emb -from .compressor import compressor_prefill_state, compressor_decode_step -from .attention import vllm_sparse_attn - - -DSV4_DEBUG_DIRECT_PREFILL_COMP = os.getenv("DSV4_DEBUG_DIRECT_PREFILL_COMP", "0") == "1" -DSV4_DEBUG_DISABLE_COMP_ATTN = os.getenv("DSV4_DEBUG_DISABLE_COMP_ATTN", "0") == "1" +from .compressor import compressor_prefill_state, compressor_decode_step, compressor_decode_step_batch +from .attention import vllm_sparse_attn_flat class DeepseekV4TransformerLayerInfer(TransformerLayerInferTpl): @@ -54,13 +49,18 @@ def __init__(self, layer_num, network_config): self.tp_q_heads = self.n_heads // self.tp_world_size_ self.tp_index_heads = self.index_n_heads // self.tp_world_size_ self.tp_groups = self.o_groups // self.tp_world_size_ + self.tp_q_head_num_ = self.tp_q_heads + self.tp_k_head_num_ = 1 + self.tp_v_head_num_ = 1 + self.tp_o_head_num_ = self.tp_q_heads + self.head_dim_ = self.head_dim self.embed_dim_ = self.hc_mult * self.hidden self.enable_ep_moe = get_env_start_args().enable_ep_moe self.indexer_score_scale = self.index_head_dim ** -0.5 self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 # ------------------------------------------------------------------ forward (HC-wrapped) - def _hc_block(self, streams, infer_state, lw, attn_fn): + def _hc_forward(self, streams, infer_state, lw, attn_forward): residual = streams collapsed, post, comb = hc_pre( streams, @@ -72,7 +72,7 @@ def _hc_block(self, streams, infer_state, lw, attn_fn): self.hc_eps, self.sinkhorn_iters, ) - o = attn_fn(lw.attn_norm_(collapsed, eps=self.eps_), infer_state, lw) + o = attn_forward(self._att_norm(collapsed, infer_state, lw), infer_state, lw) streams = hc_post(o, residual, post, comb, self.hc_mult, self.hidden) residual = streams @@ -86,17 +86,29 @@ def _hc_block(self, streams, infer_state, lw, attn_fn): self.hc_eps, self.sinkhorn_iters, ) - f = self._moe_ffn(lw.ffn_norm_(collapsed, eps=self.eps_), infer_state, lw) + f = self._ffn(self._ffn_norm(collapsed, infer_state, lw), infer_state, lw) return hc_post(f, residual, post, comb, self.hc_mult, self.hidden) def context_forward(self, streams, infer_state, lw): - return self._hc_block(streams, infer_state, lw, self._attention_prefill) + return self._hc_forward(streams, infer_state, lw, self.context_attention_forward) def token_forward(self, streams, infer_state, lw): - return self._hc_block(streams, infer_state, lw, self._attention_decode) + return self._hc_forward(streams, infer_state, lw, self.token_attention_forward) + + def _att_norm(self, x, infer_state, lw): + return lw.attn_norm_(x, eps=self.eps_) + + def _ffn_norm(self, x, infer_state, lw): + return lw.ffn_norm_(x, eps=self.eps_) - # ------------------------------------------------------------------ shared projections - def _qkv(self, x, cos_tok, sin_tok, lw): + # ------------------------------------------------------------------ shared projections / cache + def _select_rope(self, infer_state): + if self.compress_ratio: + return infer_state.position_cos_compress, infer_state.position_sin_compress + return infer_state.position_cos_sliding, infer_state.position_sin_sliding + + def _get_qkv(self, x, infer_state, lw): + cos_tok, sin_tok = self._select_rope(infer_state) T = x.shape[0] qa = lw.q_norm_(lw.wq_a_.mm(x), eps=self.eps_) q = lw.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim).float() @@ -116,10 +128,10 @@ def _qkv(self, x, cos_tok, sin_tok, lw): ], dim=1, ) - return q, kv, qa + return q, kv, qa, cos_tok, sin_tok - def _out_proj(self, o, infer_state, lw): - # o: [T, tp_q_heads, head_dim] -> inverse rope -> grouped low-rank O -> [T, hidden] + def _get_o(self, o, infer_state, lw): + # o: [T, tp_q_heads, head_dim] after inverse rope -> grouped low-rank O -> [T, hidden] T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] o = lw.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] @@ -142,22 +154,36 @@ def _inv_rope(self, o, cos_tok, sin_tok): dim=-1, ) - def _post_dense_kv(self, infer_state, req, start_pos, mem_index, kv): + def _post_cache_kv(self, cache_kv, infer_state, lw, req_idx=None, start_pos=None, mem_index=None): + if req_idx is None or start_pos is None or mem_index is None: + raise RuntimeError("DeepSeek-V4 cache write requires req_idx, start_pos, and mem_index") positions = torch.arange( start_pos, - start_pos + kv.shape[0], + start_pos + cache_kv.shape[0], device=mem_index.device, dtype=torch.long, ) infer_state.mem_manager.pack_mla_kv_to_cache( layer_index=self.layer_num_, mem_index=mem_index, - kv=kv.reshape(kv.shape[0], 1, kv.shape[-1]), - req_idx=req, + kv=cache_kv.reshape(cache_kv.shape[0], 1, cache_kv.shape[-1]), + req_idx=req_idx, positions=positions, ) return + def _get_compressor_state(self, infer_state, req): + cstate_kv, cstate_score = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) + state = { + "cstate_kv": cstate_kv, + "cstate_score": cstate_score, + } + if self.compress_ratio == 4: + idx_state = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) + state["idx_cstate_kv"] = idx_state[req, 0] + state["idx_cstate_score"] = idx_state[req, 1] + return state + def _write_compressed_kv(self, infer_state, req, entry_start, comp): slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, comp.shape[0]) if comp.shape[0] == 0: @@ -192,20 +218,61 @@ def _c4_indexer_k_from_cache(self, infer_state, req, ncomp): slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() return infer_state.mem_manager.gather_c4_indexer_k(self.layer_num_, slots) + def _run_sparse_attention_batch(self, q_chunks, kv_chunks, index_chunks, sink): + q_flat = torch.cat(q_chunks, dim=0) + kv_flat = torch.cat(kv_chunks, dim=0) + max_topk = max(t.shape[-1] for t in index_chunks) + topk = torch.full( + (q_flat.shape[0], max_topk), + -1, + dtype=torch.int32, + device=q_flat.device, + ) + offset = 0 + for idx in index_chunks: + rows = idx.shape[0] + topk[offset : offset + rows, : idx.shape[1]] = idx.to(torch.int32) + offset += rows + return vllm_sparse_attn_flat(q_flat, kv_flat, sink, topk, self.softmax_scale) + # ------------------------------------------------------------------ attention (prefill) - def _attention_prefill(self, x, infer_state, lw): + def context_attention_forward(self, x, infer_state, lw): + q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) + o = self._context_attention_wrapper_run(q, cache_kv, q_lora, x, infer_state, lw) + return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + + def _context_attention_wrapper_run(self, q, cache_kv, q_lora, x, infer_state, lw): + if torch.cuda.is_current_stream_capturing(): + q = q.contiguous() + cache_kv = cache_kv.contiguous() + q_lora = q_lora.contiguous() + x = x.contiguous() + _q = tensor_to_no_ref_tensor(q) + _cache_kv = tensor_to_no_ref_tensor(cache_kv) + _q_lora = tensor_to_no_ref_tensor(q_lora) + _x = tensor_to_no_ref_tensor(x) + + pre_capture_graph = infer_state.prefill_cuda_graph_get_current_capture_graph() + pre_capture_graph.__exit__(None, None, None) + + infer_state.prefill_cuda_graph_create_graph_obj() + infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() + o = torch.empty((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) + _o = tensor_to_no_ref_tensor(o) + + def att_func(new_infer_state): + tmp_o = self._context_attention_kernel(_q, _cache_kv, _q_lora, _x, new_infer_state, lw) + assert tmp_o.shape == _o.shape + _o.copy_(tmp_o) + return + + infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) + return o + + return self._context_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) + + def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state, lw): T = x.shape[0] - if self.compress_ratio: - cos_tok, sin_tok = ( - infer_state.position_cos_compress, - infer_state.position_sin_compress, - ) - else: - cos_tok, sin_tok = ( - infer_state.position_cos_sliding, - infer_state.position_sin_sliding, - ) - q, kv, qa = self._qkv(x, cos_tok, sin_tok, lw) sink = lw.attn_sink_.weight o = x.new_empty(T, self.tp_q_heads, self.head_dim) b_req = infer_state.b_req_idx.tolist() @@ -214,17 +281,28 @@ def _attention_prefill(self, x, infer_state, lw): ready_lens = infer_state.b_ready_cache_len.tolist() idx_q, idx_weight = self._indexer_q_weight( x, - qa, + q_lora, infer_state.position_cos_compress, infer_state.position_sin_compress, lw, ) + q_chunks = [] + kv_chunks = [] + index_chunks = [] + out_ranges = [] + kv_offset = 0 + hold_req = infer_state.req_manager.HOLD_REQUEST_ID for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): - q_r, kv_r, x_r = q[st : st + ln], kv[st : st + ln], x[st : st + ln] + if req == hold_req: + o[st : st + ln].zero_() + continue + q_r = q[st : st + ln] + cache_kv_r = cache_kv[st : st + ln] + x_r = x[st : st + ln] idx_q_r = None if idx_q is None else idx_q[st : st + ln] idx_weight_r = None if idx_weight is None else idx_weight[st : st + ln] kv_all, dense_base, n_window, ncomp, idx_comp = self._gather_prefill( - x_r, kv_r, req, ready_len, lw, infer_state + x_r, cache_kv_r, req, ready_len, lw, infer_state ) ti = self._topk_idxs_prefill( ln, @@ -237,16 +315,28 @@ def _attention_prefill(self, x, infer_state, lw): idx_comp, idx_weight_r, infer_state, - ) - o[st : st + ln] = vllm_sparse_attn(q_r.unsqueeze(0), kv_all.unsqueeze(0), sink, ti, self.softmax_scale)[0] - self._post_dense_kv( + )[0] + ti = torch.where(ti >= 0, ti + kv_offset, ti).to(torch.int32) + q_chunks.append(q_r) + kv_chunks.append(kv_all) + index_chunks.append(ti) + out_ranges.append((st, ln)) + kv_offset += kv_all.shape[0] + self._post_cache_kv( + cache_kv_r, infer_state, - req, - ready_len, - infer_state.mem_index[st : st + ln], - kv_r, + lw, + req_idx=req, + start_pos=ready_len, + mem_index=infer_state.mem_index[st : st + ln], ) - return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + if q_chunks: + attn_out = self._run_sparse_attention_batch(q_chunks, kv_chunks, index_chunks, sink) + out_offset = 0 + for st, ln in out_ranges: + o[st : st + ln] = attn_out[out_offset : out_offset + ln] + out_offset += ln + return o def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): ln = kv_r.shape[0] @@ -271,16 +361,9 @@ def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): state_pool=cstate_pool, ) comp_slots = self._write_compressed_kv(infer_state, req, 0, comp) - ( - cstate_kv, - cstate_score, - ) = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) + cstate_kv, cstate_score = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) cstate_kv.copy_(ks) cstate_score.copy_(ss) - state = { - "cstate_kv": cstate_kv, - "cstate_score": cstate_score, - } if self.compress_ratio == 4: idx_cstate_pool = infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) idx_comp, idx_ks, idx_ss, idx_cstate_pool = compressor_prefill_state( @@ -304,35 +387,15 @@ def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): idx_cstate_score = idx_state[req, 1] idx_cstate_kv.copy_(idx_ks) idx_cstate_score.copy_(idx_ss) - state.update( - { - "idx_cstate_kv": idx_cstate_kv, - "idx_cstate_score": idx_cstate_score, - } - ) - infer_state.req_manager.set_runtime_state( - req, - self.layer_num_, - state, - ) ncomp = comp.shape[0] - if DSV4_DEBUG_DISABLE_COMP_ATTN: - return kv_r, 0, ln, 0, None - if not DSV4_DEBUG_DIRECT_PREFILL_COMP: - comp = self._compressed_kv_from_cache(infer_state, req, ncomp) - idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) + comp = self._compressed_kv_from_cache(infer_state, req, ncomp) + idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) return torch.cat([kv_r, comp], dim=0), 0, ln, ncomp, idx_comp return kv_r, 0, ln, 0, None def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state): if self.compress_ratio: - try: - state = infer_state.req_manager.get_runtime_state(req, self.layer_num_) - except KeyError as exc: - raise RuntimeError( - "DeepSeek-V4 prefill chunk is missing runtime state; radix prompt cache " - "must stay disabled until V4 managed token cache is implemented." - ) from exc + state = self._get_compressor_state(infer_state, req) cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) idx_cstate_pool = ( infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) @@ -392,8 +455,6 @@ def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state): dense = torch.cat([cached_dense, kv_r], dim=0) comp = self._compressed_kv_from_cache(infer_state, req, ncomp) idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) - if DSV4_DEBUG_DISABLE_COMP_ATTN: - return dense, dense_base, dense.shape[0], 0, None return ( torch.cat([dense, comp], dim=0), dense_base, @@ -444,23 +505,201 @@ def _topk_idxs_prefill( return torch.cat([win, comp], dim=1).int().unsqueeze(0) return win.int().unsqueeze(0) + def _decode_dense_kv_graph(self, infer_state): + req = infer_state.b_req_idx.long() + seq = infer_state.b_seq_len.long() + B = req.shape[0] + device = infer_state.b_seq_len.device + offsets = torch.arange(self.window, device=device, dtype=torch.long) + win_len = torch.minimum(seq, torch.full_like(seq, self.window)) + start = seq - win_len + pos = start.unsqueeze(1) + offsets.unsqueeze(0) + valid = offsets.unsqueeze(0) < win_len.unsqueeze(1) + hold = infer_state.mem_manager.swa_pool.HOLD_TOKEN_MEMINDEX + safe_pos = torch.where(valid, pos, torch.zeros_like(pos)).long() + full_slots = infer_state.req_manager.req_to_token_indexs[req.unsqueeze(1), safe_pos].long() + swa_slots = infer_state.mem_manager.full_to_swa_indexs[full_slots].long() + slot_valid = valid & (swa_slots >= 0) + swa_slots = torch.where(slot_valid, swa_slots, torch.full_like(swa_slots, hold)) + kv = infer_state.mem_manager.gather_mla_kv_from_swa_slots(self.layer_num_, swa_slots.reshape(-1)) + return kv.view(B, self.window, self.head_dim), valid + + def _decode_all_compressed_kv_graph(self, infer_state, ratio): + req = infer_state.b_req_idx.long() + seq = infer_state.b_seq_len.long() + B = req.shape[0] + device = infer_state.b_seq_len.device + max_comp = max(1, infer_state.max_kv_seq_len // ratio) + offsets = torch.arange(max_comp, device=device, dtype=torch.long) + ncomp = torch.div(seq, ratio, rounding_mode="floor") + valid = offsets.unsqueeze(0) < ncomp.unsqueeze(1) + safe_offsets = torch.where(valid, offsets.unsqueeze(0), torch.zeros_like(offsets).unsqueeze(0)) + if ratio == 4: + table = infer_state.req_manager.req_to_c4_indexs + hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX + else: + table = infer_state.req_manager.req_to_c128_indexs + hold = infer_state.mem_manager.c128_pool.HOLD_TOKEN_MEMINDEX + slots = table[req.unsqueeze(1), safe_offsets].long() + slots = torch.where(valid, slots, torch.full_like(slots, hold)) + kv = infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots.reshape(-1)) + kv = kv.view(B, max_comp, self.head_dim) + if ratio != 4: + return kv, None, valid, ncomp + idx_k = infer_state.mem_manager.gather_c4_indexer_k(self.layer_num_, slots.reshape(-1)) + idx_k = idx_k.view(B, max_comp, self.index_head_dim) + return kv, idx_k, valid, ncomp + + def _decode_c4_topk_graph(self, idx_q, idx_weight, idx_comp, valid_comp, ncomp, infer_state): + scores = torch.einsum("bhd,bnd->bhn", idx_q.float(), idx_comp.float()) + scores = F.relu(scores) * self.indexer_score_scale + index_scores = (scores * idx_weight.unsqueeze(-1)).sum(dim=1) + if self.tp_world_size_ > 1: + all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + index_scores = index_scores.masked_fill(~valid_comp, float("-inf")) + top = index_scores.topk(self.index_topk, dim=-1).indices + valid = top < ncomp.unsqueeze(1) + return torch.where(valid, top, torch.zeros_like(top)), valid + + def _decode_compressed_candidates_graph(self, idx_q, idx_weight, infer_state): + if self.compress_ratio == 4: + _, idx_comp, valid_all, ncomp = self._decode_all_compressed_kv_graph(infer_state, 4) + top, valid = self._decode_c4_topk_graph(idx_q, idx_weight, idx_comp, valid_all, ncomp, infer_state) + req = infer_state.b_req_idx.long() + slots = infer_state.req_manager.req_to_c4_indexs[req.unsqueeze(1), top].long() + hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX + slots = torch.where(valid, slots, torch.full_like(slots, hold)) + comp = infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots.reshape(-1)) + return comp.view(req.shape[0], self.index_topk, self.head_dim), valid + comp, _, valid, _ = self._decode_all_compressed_kv_graph(infer_state, 128) + return comp, valid + + def _write_decode_compressed_entry_graph(self, x, infer_state, lw, ratio): + req = infer_state.b_req_idx + start_pos = infer_state.b_seq_len.long() - 1 + if ratio == 4: + state_all = infer_state.req_manager.get_c4_compress_state(self.layer_num_) + table = infer_state.req_manager.req_to_c4_indexs + hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX + else: + state_all = infer_state.req_manager.get_c128_compress_state(self.layer_num_) + table = infer_state.req_manager.req_to_c128_indexs + hold = infer_state.mem_manager.c128_pool.HOLD_TOKEN_MEMINDEX + + entry, should = compressor_decode_step_batch( + x, + lw.compressor_wkv_.mm_param.weight, + lw.compressor_wgate_.mm_param.weight, + lw.compressor_norm_.weight, + lw.compressor_ape_.weight, + ratio, + self.head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + state_all, + req, + start_pos, + ) + entry_idx = torch.clamp(torch.div(infer_state.b_seq_len.long(), ratio, rounding_mode="floor") - 1, min=0) + slots = table[req.long(), entry_idx].long() + slots = torch.where(should, slots, torch.full_like(slots, hold)) + infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, entry) + + if ratio == 4: + idx_state_all = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) + idx_entry, idx_should = compressor_decode_step_batch( + x, + lw.idx_cmp_wkv_.mm_param.weight, + lw.idx_cmp_wgate_.mm_param.weight, + lw.idx_cmp_norm_.weight, + lw.idx_cmp_ape_.weight, + 4, + self.index_head_dim, + self.rope_dim, + infer_state.cos_compress_table, + infer_state.sin_compress_table, + self.eps_, + idx_state_all, + req, + start_pos, + ) + idx_slots = torch.where(idx_should, slots, torch.full_like(slots, hold)) + infer_state.mem_manager.pack_c4_indexer_k_to_cache(self.layer_num_, idx_slots, idx_entry) + return + # ------------------------------------------------------------------ attention (decode) - def _attention_decode(self, x, infer_state, lw): - B = x.shape[0] # one new token per request + def token_attention_forward(self, x, infer_state, lw): + q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) + if infer_state.is_cuda_graph: + o = self._token_attention_kernel_cuda_graph(q, cache_kv, q_lora, x, infer_state, lw) + else: + o = self._token_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) + return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + + def _token_attention_kernel_cuda_graph(self, q, cache_kv, q_lora, x, infer_state, lw): + sink = lw.attn_sink_.weight + infer_state.mem_manager.pack_decode_mla_kv_to_cache( + self.layer_num_, + infer_state.b_req_idx, + infer_state.b_seq_len, + infer_state.mem_index, + cache_kv.reshape(cache_kv.shape[0], 1, cache_kv.shape[-1]), + ) + idx_q, idx_weight = self._indexer_q_weight( + x, + q_lora, + infer_state.position_cos_compress, + infer_state.position_sin_compress, + lw, + ) if self.compress_ratio: - cos_tok, sin_tok = ( - infer_state.position_cos_compress, - infer_state.position_sin_compress, - ) + self._write_decode_compressed_entry_graph(x, infer_state, lw, self.compress_ratio) + + dense_kv, dense_valid = self._decode_dense_kv_graph(infer_state) + B = q.shape[0] + device = q.device + if self.compress_ratio: + comp_kv, comp_valid = self._decode_compressed_candidates_graph(idx_q, idx_weight, infer_state) + kv_all = torch.cat([dense_kv, comp_kv], dim=1) + comp_offsets = torch.arange(comp_kv.shape[1], device=device, dtype=torch.int32) else: - cos_tok, sin_tok = ( - infer_state.position_cos_sliding, - infer_state.position_sin_sliding, + kv_all = dense_kv + comp_valid = None + comp_offsets = None + + total_k = kv_all.shape[1] + base = torch.arange(B, device=device, dtype=torch.int32).unsqueeze(1) * total_k + dense_offsets = torch.arange(self.window, device=device, dtype=torch.int32) + dense_topk = torch.where( + dense_valid, + base + dense_offsets.unsqueeze(0), + torch.full((B, self.window), -1, device=device, dtype=torch.int32), + ) + if self.compress_ratio: + comp_topk = torch.where( + comp_valid, + base + self.window + comp_offsets.unsqueeze(0), + torch.full((B, comp_kv.shape[1]), -1, device=device, dtype=torch.int32), ) - q, kv, qa = self._qkv(x, cos_tok, sin_tok, lw) # [B, heads, hd], [B, hd] + topk = torch.cat([dense_topk, comp_topk], dim=1) + else: + topk = dense_topk + return vllm_sparse_attn_flat( + q, + kv_all.reshape(-1, self.head_dim), + sink, + topk, + self.softmax_scale, + already_compact=True, + ) + + def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state, lw): + B = x.shape[0] # one new token per request idx_q, idx_weight = self._indexer_q_weight( x, - qa, + q_lora, infer_state.position_cos_compress, infer_state.position_sin_compress, lw, @@ -469,23 +708,27 @@ def _attention_decode(self, x, infer_state, lw): b_req = infer_state.b_req_idx.tolist() seqlens = infer_state.b_seq_len.tolist() o = x.new_empty(B, self.tp_q_heads, self.head_dim) + hold_req = infer_state.req_manager.HOLD_REQUEST_ID + q_chunks = [] + kv_chunks = [] + index_chunks = [] + out_rows = [] + kv_offset = 0 for i, (req, seq) in enumerate(zip(b_req, seqlens)): + if req == hold_req: + o[i].zero_() + continue start_pos = seq - 1 - self._post_dense_kv( + self._post_cache_kv( + cache_kv[i : i + 1], infer_state, - req, - start_pos, - infer_state.mem_index[i : i + 1], - kv[i : i + 1], + lw, + req_idx=req, + start_pos=start_pos, + mem_index=infer_state.mem_index[i : i + 1], ) if self.compress_ratio: - try: - stt = infer_state.req_manager.get_runtime_state(req, self.layer_num_) - except KeyError as exc: - raise RuntimeError( - "DeepSeek-V4 decode is missing runtime state; radix prompt cache " - "must stay disabled until V4 managed token cache is implemented." - ) from exc + stt = self._get_compressor_state(infer_state, req) cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) e = compressor_decode_step( x[i], @@ -538,12 +781,7 @@ def _attention_decode(self, x, infer_state, lw): win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) comp_kv = self._compressed_kv_from_cache(infer_state, req, seq // self.compress_ratio) idx_comp = self._c4_indexer_k_from_cache(infer_state, req, comp_kv.shape[0]) - if DSV4_DEBUG_DISABLE_COMP_ATTN: - comp_kv = None - idx_comp = None - kv_all = win_kv - else: - kv_all = torch.cat([win_kv, comp_kv], dim=0) + kv_all = torch.cat([win_kv, comp_kv], dim=0) else: win_start = max(0, seq - self.window) win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) @@ -559,15 +797,18 @@ def _attention_decode(self, x, infer_state, lw): seq, x.device, infer_state, - ) - o[i] = vllm_sparse_attn( - q[i].view(1, 1, self.tp_q_heads, self.head_dim), - kv_all.unsqueeze(0), - sink, - ti, - self.softmax_scale, )[0, 0] - return self._out_proj(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) + ti = torch.where(ti >= 0, ti + kv_offset, ti).view(1, -1).to(torch.int32) + q_chunks.append(q[i : i + 1]) + kv_chunks.append(kv_all) + index_chunks.append(ti) + out_rows.append(i) + kv_offset += kv_all.shape[0] + if q_chunks: + attn_out = self._run_sparse_attention_batch(q_chunks, kv_chunks, index_chunks, sink) + for row, row_out in zip(out_rows, attn_out): + o[row] = row_out + return o def _indexer_q_weight(self, x, qa, cos_tok, sin_tok, lw): if self.compress_ratio != 4: @@ -705,7 +946,7 @@ def _fp4_experts_marlin(self, x, weights, indices, experts): clamp_limit=float(self.swiglu_limit), ) - def _moe_ffn(self, x, infer_state, lw): + def _ffn(self, x, infer_state, lw): gw = lw.gate_weight_.mm_param.weight logits = F.linear(x.float(), gw.float()).contiguous() weights, indices = self._select_experts(logits, infer_state, lw) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index c87f2fdebd..915e45b9c9 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -37,12 +37,13 @@ from lightllm.distributed.communication_op import dist_group_manager logger = init_logger(__name__) +DSV4_DECODE_CUDAGRAPH_MAX_LEN = 8192 class DeepseekV4DirectSparseAttBackend(BaseAttBackend): """Lifecycle placeholder for V4 direct attention. - V4 attention is currently driven inside the layer by `vllm_sparse_attn()`, not by the generic + V4 attention is currently driven inside the layer, not by the generic `infer_state.prefill_att_state.prefill_att()` / `decode_att()` backend selector. """ @@ -58,7 +59,7 @@ def init_state(self): return def prefill_att(self, *args, **kwargs): - raise RuntimeError("DeepSeek-V4 attention is executed directly by vllm_sparse_attn() in layer_infer.") + raise RuntimeError("DeepSeek-V4 attention is executed directly in layer_infer.") class DeepseekV4DirectSparseDecodeAttState(BaseDecodeAttState): @@ -66,7 +67,7 @@ def init_state(self): return def decode_att(self, *args, **kwargs): - raise RuntimeError("DeepSeek-V4 attention is executed directly by vllm_sparse_attn() in layer_infer.") + raise RuntimeError("DeepSeek-V4 attention is executed directly in layer_infer.") @ModelRegistry("deepseek_v4") @@ -79,6 +80,7 @@ class DeepseekV4TpPartModel(LlamaTpPartModel): transformer_layer_infer_class = DeepseekV4TransformerLayerInfer infer_state_class = DeepseekV4InferStateInfo + _logged_prefill_graph_prefix_skip = False def _verify_params(self): assert self.load_way == "HF", "only support HF format weights" @@ -137,6 +139,28 @@ def _init_mem_manager(self): self.req_manager.bind_mem_manager(self.mem_manager) return + def _init_cudagraph(self): + if not self.disable_cudagraph and self.graph_max_len_in_batch > DSV4_DECODE_CUDAGRAPH_MAX_LEN: + logger.info( + "DeepSeek-V4 caps decode cudagraph max_len_in_batch from %s to %s for the current " + "graph-safe sparse-attention fallback; longer decode batches run eager.", + self.graph_max_len_in_batch, + DSV4_DECODE_CUDAGRAPH_MAX_LEN, + ) + self.graph_max_len_in_batch = DSV4_DECODE_CUDAGRAPH_MAX_LEN + return super()._init_cudagraph() + + def _can_run_prefill_cudagraph(self, infer_state, handle_token_num): + if infer_state.prefix_total_token_num == 0: + return True + if not self._logged_prefill_graph_prefix_skip: + logger.info( + "DeepSeek-V4 skips prefill cudagraph for prompt-cache extension batches; " + "no-prefix prefill batches still use prefill cudagraph." + ) + self._logged_prefill_graph_prefix_skip = True + return False + def _init_att_backend(self): self.prefill_att_backend = DeepseekV4DirectSparseAttBackend(model=self) self.decode_att_backend = DeepseekV4DirectSparseAttBackend(model=self) From 19866d02ab1a59afc6a65deccd86d05eb9d93459 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 8 Jun 2026 01:57:00 +0000 Subject: [PATCH 005/214] refact tokenizer --- lightllm/models/deepseek3_2/model.py | 71 +++++++++++++++------------- lightllm/models/deepseek_v4/model.py | 61 ++++++++++++++++++++++-- 2 files changed, 95 insertions(+), 37 deletions(-) diff --git a/lightllm/models/deepseek3_2/model.py b/lightllm/models/deepseek3_2/model.py index cd33386666..5831044311 100644 --- a/lightllm/models/deepseek3_2/model.py +++ b/lightllm/models/deepseek3_2/model.py @@ -1,20 +1,14 @@ import copy from lightllm.models.registry import ModelRegistry from lightllm.models.deepseek2.model import Deepseek2TpPartModel -from lightllm.models.deepseek3_2.layer_weights.transformer_layer_weight import ( - Deepseek3_2TransformerLayerWeight, -) -from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import ( - Deepseek3_2TransformerLayerInfer, -) -from lightllm.common.basemodel.attention import ( - get_nsa_prefill_att_backend_class, - get_nsa_decode_att_backend_class, -) +from lightllm.models.deepseek3_2.layer_weights.transformer_layer_weight import Deepseek3_2TransformerLayerWeight +from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer +from lightllm.common.basemodel.attention import get_nsa_prefill_att_backend_class, get_nsa_decode_att_backend_class @ModelRegistry(["deepseek_v32"]) class Deepseek3_2TpPartModel(Deepseek2TpPartModel): + # weight class transformer_weight_class = Deepseek3_2TransformerLayerWeight @@ -27,11 +21,24 @@ def _init_att_backend(self): return -class DeepSeekChatTokenizerBase: +class DeepSeekV32Tokenizer: + """Tokenizer wrapper for DeepSeek-V3.2 that uses the Python-based + encoding_dsv32 module instead of Jinja chat templates. + + DeepSeek-V3.2's tokenizer_config.json does not ship with a Jinja chat + template, so ``apply_chat_template`` would fail without either a manually + supplied ``--chat_template`` file or this wrapper. + """ + def __init__(self, tokenizer): self.tokenizer = tokenizer + # Cache added vocabulary for performance (HuggingFace can be slow). self._added_vocab = None + # ------------------------------------------------------------------ + # Attribute delegation – everything not overridden goes to the inner + # tokenizer so that encode/decode/vocab_size/eos_token_id/… all work. + # ------------------------------------------------------------------ def __getattr__(self, name): return getattr(self.tokenizer, name) @@ -40,9 +47,9 @@ def get_added_vocab(self): self._added_vocab = self.tokenizer.get_added_vocab() return self._added_vocab - def _encode_messages(self, msgs, thinking_mode, kwargs): - raise NotImplementedError("subclass must provide DeepSeek encode_messages") - + # ------------------------------------------------------------------ + # Core override: route apply_chat_template through encode_messages. + # ------------------------------------------------------------------ def apply_chat_template( self, conversation=None, @@ -51,16 +58,27 @@ def apply_chat_template( tokenize=False, add_generation_prompt=True, thinking=None, - enable_thinking=None, **kwargs, ): + from lightllm.models.deepseek3_2.encoding_dsv32 import encode_messages, render_tools + msgs = conversation if conversation is not None else messages if msgs is None: raise ValueError("Either 'conversation' or 'messages' must be provided") + # Deep copy to avoid mutating the caller's messages. msgs = copy.deepcopy(msgs) + # Determine thinking mode. + thinking_mode = "thinking" if thinking else "chat" + + # Inject tools into the first system message (or create one) so that + # encode_messages / render_message picks them up. if tools: + # build_prompt passes tools as bare function dicts: + # [{"name": "f", "description": "...", "parameters": {...}}] + # encoding_dsv32's render_message expects OpenAI wrapper format: + # [{"type": "function", "function": {...}}] wrapped_tools = [] for t in tools: if "function" in t: @@ -77,27 +95,16 @@ def apply_chat_template( break if not injected: + # Prepend a system message that carries the tools. msgs.insert(0, {"role": "system", "content": "", "tools": wrapped_tools}) - if thinking is None: - thinking = bool(enable_thinking) if enable_thinking is not None else False - thinking_mode = "thinking" if thinking else "chat" - prompt = self._encode_messages(msgs, thinking_mode, kwargs) - - if tokenize: - return self.tokenizer.encode(prompt, add_special_tokens=False) - return prompt - - -class DeepSeekV32Tokenizer(DeepSeekChatTokenizerBase): - """Tokenizer wrapper for DeepSeek-V3.2's Python-based encoding_dsv32 module.""" - - def _encode_messages(self, msgs, thinking_mode, kwargs): - from lightllm.models.deepseek3_2.encoding_dsv32 import encode_messages - - return encode_messages( + prompt = encode_messages( msgs, thinking_mode=thinking_mode, drop_thinking=kwargs.get("drop_thinking", True), add_default_bos_token=kwargs.get("add_default_bos_token", True), ) + + if tokenize: + return self.tokenizer.encode(prompt, add_special_tokens=False) + return prompt diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 915e45b9c9..914804fe86 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -1,3 +1,4 @@ +import copy import importlib.util import os @@ -27,7 +28,6 @@ DeepseekV4TransformerLayerInfer, ) from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo -from lightllm.models.deepseek3_2.model import DeepSeekChatTokenizerBase from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, linear_ramp_mask, @@ -212,13 +212,22 @@ def build(base, factor, orig_max): return -class DeepSeekV4Tokenizer(DeepSeekChatTokenizerBase): +class DeepSeekV4Tokenizer: """Tokenizer wrapper for DeepSeek-V4's Python prompt encoding.""" def __init__(self, tokenizer, model_dir): - super().__init__(tokenizer) + self.tokenizer = tokenizer self.model_dir = model_dir self._encoding_module = None + self._added_vocab = None + + def __getattr__(self, name): + return getattr(self.tokenizer, name) + + def get_added_vocab(self): + if self._added_vocab is None: + self._added_vocab = self.tokenizer.get_added_vocab() + return self._added_vocab def _get_encoding_module(self): if self._encoding_module is not None: @@ -236,12 +245,54 @@ def _get_encoding_module(self): self._encoding_module = module return module - def _encode_messages(self, msgs, thinking_mode, kwargs): + def apply_chat_template( + self, + conversation=None, + messages=None, + tools=None, + tokenize=False, + add_generation_prompt=True, + thinking=None, + enable_thinking=None, + **kwargs, + ): + msgs = conversation if conversation is not None else messages + if msgs is None: + raise ValueError("Either 'conversation' or 'messages' must be provided") + + msgs = copy.deepcopy(msgs) + + if tools: + wrapped_tools = [] + for tool in tools: + if "function" in tool: + wrapped_tools.append(tool) + else: + wrapped_tools.append({"type": "function", "function": tool}) + + injected = False + for msg in msgs: + if msg.get("role") == "system": + existing = msg.get("tools") or [] + msg["tools"] = existing + wrapped_tools + injected = True + break + + if not injected: + msgs.insert(0, {"role": "system", "content": "", "tools": wrapped_tools}) + + if thinking is None: + thinking = bool(enable_thinking) if enable_thinking is not None else False + thinking_mode = "thinking" if thinking else "chat" encoding = self._get_encoding_module() - return encoding.encode_messages( + prompt = encoding.encode_messages( msgs, thinking_mode=thinking_mode, drop_thinking=kwargs.get("drop_thinking", True), add_default_bos_token=kwargs.get("add_default_bos_token", True), reasoning_effort=kwargs.get("reasoning_effort"), ) + + if tokenize: + return self.tokenizer.encode(prompt, add_special_tokens=False) + return prompt From 29c6082a897485f0d4d4cbcddd1ec6b1f2f01d3b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 8 Jun 2026 04:58:53 +0000 Subject: [PATCH 006/214] add statement --- lightllm/models/deepseek_v4/infer_struct.py | 5 ++ .../layer_infer/post_layer_infer.py | 3 +- .../layer_infer/pre_layer_infer.py | 7 +- .../layer_infer/transformer_layer_infer.py | 77 ++++++++++--------- lightllm/models/deepseek_v4/model.py | 24 ++---- 5 files changed, 59 insertions(+), 57 deletions(-) diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index 6bc402cd28..d0c2745161 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -1,8 +1,13 @@ import torch from lightllm.common.basemodel import InferStateInfo +from lightllm.common.req_manager import DeepseekV4ReqManager +from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager class DeepseekV4InferStateInfo(InferStateInfo): + req_manager: DeepseekV4ReqManager + mem_manager: DeepseekV4MemoryManager + """Per-token interleaved-rope cos/sin for the two rope variants (sliding / compressed), following the gemma4 two-variant convention (_cos_cached_* -> position_cos_*). Also exposes the full compressed cos/sin tables, which the KV compressor indexes at window positions (not per-token).""" diff --git a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py index 87951e7360..c23d03afb7 100644 --- a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py @@ -1,11 +1,12 @@ from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer from .hyper_connection import hc_head +from ..infer_struct import DeepseekV4InferStateInfo class DeepseekV4PostLayerInfer(LlamaPostLayerInfer): """Collapse the hc_mult residual streams (hc_head) to [T, hidden], then final norm + lm_head.""" - def token_forward(self, input_embdings, infer_state, layer_weight): + def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): cfg = layer_weight.network_config_ collapsed = hc_head( input_embdings, diff --git a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py index 0be99ecbab..d83e3082b8 100644 --- a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py @@ -2,12 +2,13 @@ import torch.distributed as dist from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer from lightllm.distributed.communication_op import all_reduce +from ..infer_struct import DeepseekV4InferStateInfo class DeepseekV4PreLayerInfer(LlamaPreLayerInfer): """Token embedding, then expand to the hc_mult parallel residual streams [T, hc_mult*hidden].""" - def _embed_and_expand(self, input_ids, infer_state, layer_weight): + def _embed_and_expand(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): emb = layer_weight.wte_weight_(input_ids=input_ids, alloc_func=self.alloc_tensor) # [T, hidden] if self.tp_world_size_ > 1: all_reduce(emb, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) @@ -15,8 +16,8 @@ def _embed_and_expand(self, input_ids, infer_state, layer_weight): t, hidden = emb.shape return emb.unsqueeze(1).expand(t, hc_mult, hidden).reshape(t, hc_mult * hidden).contiguous() - def context_forward(self, input_ids, infer_state, layer_weight): + def context_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): return self._embed_and_expand(input_ids, infer_state, layer_weight) - def token_forward(self, input_ids, infer_state, layer_weight): + def token_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): return self._embed_and_expand(input_ids, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 864104fd32..11209d39fd 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -7,6 +7,7 @@ from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from .hyper_connection import hc_pre, hc_post from ..triton_kernel.rotary_emb import apply_rotary_emb +from ..infer_struct import DeepseekV4InferStateInfo from .compressor import compressor_prefill_state, compressor_decode_step, compressor_decode_step_batch from .attention import vllm_sparse_attn_flat @@ -60,7 +61,7 @@ def __init__(self, layer_num, network_config): self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 # ------------------------------------------------------------------ forward (HC-wrapped) - def _hc_forward(self, streams, infer_state, lw, attn_forward): + def _hc_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw, attn_forward): residual = streams collapsed, post, comb = hc_pre( streams, @@ -89,25 +90,25 @@ def _hc_forward(self, streams, infer_state, lw, attn_forward): f = self._ffn(self._ffn_norm(collapsed, infer_state, lw), infer_state, lw) return hc_post(f, residual, post, comb, self.hc_mult, self.hidden) - def context_forward(self, streams, infer_state, lw): + def context_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw): return self._hc_forward(streams, infer_state, lw, self.context_attention_forward) - def token_forward(self, streams, infer_state, lw): + def token_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw): return self._hc_forward(streams, infer_state, lw, self.token_attention_forward) - def _att_norm(self, x, infer_state, lw): + def _att_norm(self, x, infer_state: DeepseekV4InferStateInfo, lw): return lw.attn_norm_(x, eps=self.eps_) - def _ffn_norm(self, x, infer_state, lw): + def _ffn_norm(self, x, infer_state: DeepseekV4InferStateInfo, lw): return lw.ffn_norm_(x, eps=self.eps_) # ------------------------------------------------------------------ shared projections / cache - def _select_rope(self, infer_state): + def _select_rope(self, infer_state: DeepseekV4InferStateInfo): if self.compress_ratio: return infer_state.position_cos_compress, infer_state.position_sin_compress return infer_state.position_cos_sliding, infer_state.position_sin_sliding - def _get_qkv(self, x, infer_state, lw): + def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, lw): cos_tok, sin_tok = self._select_rope(infer_state) T = x.shape[0] qa = lw.q_norm_(lw.wq_a_.mm(x), eps=self.eps_) @@ -130,7 +131,7 @@ def _get_qkv(self, x, infer_state, lw): ) return q, kv, qa, cos_tok, sin_tok - def _get_o(self, o, infer_state, lw): + def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, lw): # o: [T, tp_q_heads, head_dim] after inverse rope -> grouped low-rank O -> [T, hidden] T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] @@ -154,7 +155,9 @@ def _inv_rope(self, o, cos_tok, sin_tok): dim=-1, ) - def _post_cache_kv(self, cache_kv, infer_state, lw, req_idx=None, start_pos=None, mem_index=None): + def _post_cache_kv( + self, cache_kv, infer_state: DeepseekV4InferStateInfo, lw, req_idx=None, start_pos=None, mem_index=None + ): if req_idx is None or start_pos is None or mem_index is None: raise RuntimeError("DeepSeek-V4 cache write requires req_idx, start_pos, and mem_index") positions = torch.arange( @@ -172,7 +175,7 @@ def _post_cache_kv(self, cache_kv, infer_state, lw, req_idx=None, start_pos=None ) return - def _get_compressor_state(self, infer_state, req): + def _get_compressor_state(self, infer_state: DeepseekV4InferStateInfo, req): cstate_kv, cstate_score = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) state = { "cstate_kv": cstate_kv, @@ -184,26 +187,26 @@ def _get_compressor_state(self, infer_state, req): state["idx_cstate_score"] = idx_state[req, 1] return state - def _write_compressed_kv(self, infer_state, req, entry_start, comp): + def _write_compressed_kv(self, infer_state: DeepseekV4InferStateInfo, req, entry_start, comp): slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, comp.shape[0]) if comp.shape[0] == 0: return slots infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) return slots - def _write_c4_indexer_k(self, infer_state, slots, idx_comp): + def _write_c4_indexer_k(self, infer_state: DeepseekV4InferStateInfo, slots, idx_comp): if idx_comp is None or idx_comp.shape[0] == 0: return infer_state.mem_manager.pack_c4_indexer_k_to_cache(self.layer_num_, slots, idx_comp) return - def _dense_kv_from_cache(self, infer_state, req, start_pos, end_pos): + def _dense_kv_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, start_pos, end_pos): if end_pos <= start_pos: return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") slots = infer_state.req_manager.req_to_token_indexs[req, start_pos:end_pos].long() return infer_state.mem_manager.gather_mla_kv(self.layer_num_, slots) - def _compressed_kv_from_cache(self, infer_state, req, ncomp): + def _compressed_kv_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, ncomp): if ncomp == 0: return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") if self.compress_ratio == 4: @@ -212,7 +215,7 @@ def _compressed_kv_from_cache(self, infer_state, req, ncomp): slots = infer_state.req_manager.req_to_c128_indexs[req, :ncomp].long() return infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots) - def _c4_indexer_k_from_cache(self, infer_state, req, ncomp): + def _c4_indexer_k_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, ncomp): if self.compress_ratio != 4 or ncomp == 0: return None slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() @@ -236,12 +239,12 @@ def _run_sparse_attention_batch(self, q_chunks, kv_chunks, index_chunks, sink): return vllm_sparse_attn_flat(q_flat, kv_flat, sink, topk, self.softmax_scale) # ------------------------------------------------------------------ attention (prefill) - def context_attention_forward(self, x, infer_state, lw): + def context_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, lw): q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) o = self._context_attention_wrapper_run(q, cache_kv, q_lora, x, infer_state, lw) return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) - def _context_attention_wrapper_run(self, q, cache_kv, q_lora, x, infer_state, lw): + def _context_attention_wrapper_run(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): if torch.cuda.is_current_stream_capturing(): q = q.contiguous() cache_kv = cache_kv.contiguous() @@ -260,7 +263,7 @@ def _context_attention_wrapper_run(self, q, cache_kv, q_lora, x, infer_state, lw o = torch.empty((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) _o = tensor_to_no_ref_tensor(o) - def att_func(new_infer_state): + def att_func(new_infer_state: DeepseekV4InferStateInfo): tmp_o = self._context_attention_kernel(_q, _cache_kv, _q_lora, _x, new_infer_state, lw) assert tmp_o.shape == _o.shape _o.copy_(tmp_o) @@ -271,7 +274,7 @@ def att_func(new_infer_state): return self._context_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) - def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state, lw): + def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): T = x.shape[0] sink = lw.attn_sink_.weight o = x.new_empty(T, self.tp_q_heads, self.head_dim) @@ -338,7 +341,7 @@ def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state, lw): out_offset += ln return o - def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): + def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state: DeepseekV4InferStateInfo): ln = kv_r.shape[0] idx_comp = None if ready_len > 0: @@ -393,7 +396,7 @@ def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state): return torch.cat([kv_r, comp], dim=0), 0, ln, ncomp, idx_comp return kv_r, 0, ln, 0, None - def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state): + def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state: DeepseekV4InferStateInfo): if self.compress_ratio: state = self._get_compressor_state(infer_state, req) cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) @@ -484,7 +487,7 @@ def _topk_idxs_prefill( idx_q, idx_comp, idx_weight, - infer_state, + infer_state: DeepseekV4InferStateInfo, ): t = torch.arange(seqlen, device=device) abs_pos = t + base_pos @@ -505,7 +508,7 @@ def _topk_idxs_prefill( return torch.cat([win, comp], dim=1).int().unsqueeze(0) return win.int().unsqueeze(0) - def _decode_dense_kv_graph(self, infer_state): + def _decode_dense_kv_graph(self, infer_state: DeepseekV4InferStateInfo): req = infer_state.b_req_idx.long() seq = infer_state.b_seq_len.long() B = req.shape[0] @@ -524,7 +527,7 @@ def _decode_dense_kv_graph(self, infer_state): kv = infer_state.mem_manager.gather_mla_kv_from_swa_slots(self.layer_num_, swa_slots.reshape(-1)) return kv.view(B, self.window, self.head_dim), valid - def _decode_all_compressed_kv_graph(self, infer_state, ratio): + def _decode_all_compressed_kv_graph(self, infer_state: DeepseekV4InferStateInfo, ratio): req = infer_state.b_req_idx.long() seq = infer_state.b_seq_len.long() B = req.shape[0] @@ -550,7 +553,9 @@ def _decode_all_compressed_kv_graph(self, infer_state, ratio): idx_k = idx_k.view(B, max_comp, self.index_head_dim) return kv, idx_k, valid, ncomp - def _decode_c4_topk_graph(self, idx_q, idx_weight, idx_comp, valid_comp, ncomp, infer_state): + def _decode_c4_topk_graph( + self, idx_q, idx_weight, idx_comp, valid_comp, ncomp, infer_state: DeepseekV4InferStateInfo + ): scores = torch.einsum("bhd,bnd->bhn", idx_q.float(), idx_comp.float()) scores = F.relu(scores) * self.indexer_score_scale index_scores = (scores * idx_weight.unsqueeze(-1)).sum(dim=1) @@ -561,7 +566,7 @@ def _decode_c4_topk_graph(self, idx_q, idx_weight, idx_comp, valid_comp, ncomp, valid = top < ncomp.unsqueeze(1) return torch.where(valid, top, torch.zeros_like(top)), valid - def _decode_compressed_candidates_graph(self, idx_q, idx_weight, infer_state): + def _decode_compressed_candidates_graph(self, idx_q, idx_weight, infer_state: DeepseekV4InferStateInfo): if self.compress_ratio == 4: _, idx_comp, valid_all, ncomp = self._decode_all_compressed_kv_graph(infer_state, 4) top, valid = self._decode_c4_topk_graph(idx_q, idx_weight, idx_comp, valid_all, ncomp, infer_state) @@ -574,7 +579,7 @@ def _decode_compressed_candidates_graph(self, idx_q, idx_weight, infer_state): comp, _, valid, _ = self._decode_all_compressed_kv_graph(infer_state, 128) return comp, valid - def _write_decode_compressed_entry_graph(self, x, infer_state, lw, ratio): + def _write_decode_compressed_entry_graph(self, x, infer_state: DeepseekV4InferStateInfo, lw, ratio): req = infer_state.b_req_idx start_pos = infer_state.b_seq_len.long() - 1 if ratio == 4: @@ -630,7 +635,7 @@ def _write_decode_compressed_entry_graph(self, x, infer_state, lw, ratio): return # ------------------------------------------------------------------ attention (decode) - def token_attention_forward(self, x, infer_state, lw): + def token_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, lw): q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) if infer_state.is_cuda_graph: o = self._token_attention_kernel_cuda_graph(q, cache_kv, q_lora, x, infer_state, lw) @@ -638,7 +643,7 @@ def token_attention_forward(self, x, infer_state, lw): o = self._token_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) - def _token_attention_kernel_cuda_graph(self, q, cache_kv, q_lora, x, infer_state, lw): + def _token_attention_kernel_cuda_graph(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): sink = lw.attn_sink_.weight infer_state.mem_manager.pack_decode_mla_kv_to_cache( self.layer_num_, @@ -695,7 +700,7 @@ def _token_attention_kernel_cuda_graph(self, q, cache_kv, q_lora, x, infer_state already_compact=True, ) - def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state, lw): + def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): B = x.shape[0] # one new token per request idx_q, idx_weight = self._indexer_q_weight( x, @@ -828,7 +833,9 @@ def _indexer_q_weight(self, x, qa, cos_tok, sin_tok, lw): idx_weight = lw.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale return idx_q, idx_weight - def _indexer_topk(self, idx_q, idx_comp, idx_weight, positions_1based, offset, infer_state): + def _indexer_topk( + self, idx_q, idx_comp, idx_weight, positions_1based, offset, infer_state: DeepseekV4InferStateInfo + ): ncomp = idx_comp.shape[0] k = min(self.index_topk, ncomp) if k == 0: @@ -896,7 +903,7 @@ def _topk_idxs_decode( idx_weight, seq_len, device, - infer_state, + infer_state: DeepseekV4InferStateInfo, ): win = torch.arange(win_len, device=device, dtype=torch.long) if comp_kv is None or comp_kv.shape[0] == 0: @@ -946,7 +953,7 @@ def _fp4_experts_marlin(self, x, weights, indices, experts): clamp_limit=float(self.swiglu_limit), ) - def _ffn(self, x, infer_state, lw): + def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, lw): gw = lw.gate_weight_.mm_param.weight logits = F.linear(x.float(), gw.float()).contiguous() weights, indices = self._select_experts(logits, infer_state, lw) @@ -968,10 +975,10 @@ def _ffn(self, x, infer_state, lw): all_reduce(out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) return out - def _select_experts(self, logits, infer_state, lw): + def _select_experts(self, logits, infer_state: DeepseekV4InferStateInfo, lw): return self._select_experts_vllm(logits, infer_state, lw) - def _select_experts_vllm(self, logits, infer_state, lw): + def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, lw): from vllm import _custom_ops as ops M = logits.shape[0] diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 914804fe86..687d5f46f0 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -47,10 +47,10 @@ class DeepseekV4DirectSparseAttBackend(BaseAttBackend): `infer_state.prefill_att_state.prefill_att()` / `decode_att()` backend selector. """ - def create_att_prefill_state(self, infer_state): + def create_att_prefill_state(self, infer_state: DeepseekV4InferStateInfo): return DeepseekV4DirectSparsePrefillAttState(backend=self, infer_state=infer_state) - def create_att_decode_state(self, infer_state): + def create_att_decode_state(self, infer_state: DeepseekV4InferStateInfo): return DeepseekV4DirectSparseDecodeAttState(backend=self, infer_state=infer_state) @@ -72,6 +72,9 @@ def decode_att(self, *args, **kwargs): @ModelRegistry("deepseek_v4") class DeepseekV4TpPartModel(LlamaTpPartModel): + req_manager: DeepseekV4ReqManager + mem_manager: DeepseekV4MemoryManager + pre_and_post_weight_class = DeepseekV4PreAndPostLayerWeight transformer_weight_class = DeepseekV4TransformerLayerWeight @@ -80,7 +83,6 @@ class DeepseekV4TpPartModel(LlamaTpPartModel): transformer_layer_infer_class = DeepseekV4TransformerLayerInfer infer_state_class = DeepseekV4InferStateInfo - _logged_prefill_graph_prefix_skip = False def _verify_params(self): assert self.load_way == "HF", "only support HF format weights" @@ -89,11 +91,6 @@ def _verify_params(self): assert self.config["index_n_heads"] % self.tp_world_size_ == 0 return - def _init_some_value(self): - super()._init_some_value() - self.head_dim_ = self.config["head_dim"] - return - def _init_req_manager(self): create_max_seq_len = 0 if self.batch_max_tokens is not None: @@ -115,9 +112,6 @@ def _init_req_manager(self): def _get_compress_rates(self, layer_num): rates = list(self.config["compress_ratios"]) - assert ( - len(rates) >= layer_num - ), f"DeepSeek-V4 compress_ratios length {len(rates)} is shorter than layer_num {layer_num}" return rates[:layer_num] def _init_mem_manager(self): @@ -150,15 +144,9 @@ def _init_cudagraph(self): self.graph_max_len_in_batch = DSV4_DECODE_CUDAGRAPH_MAX_LEN return super()._init_cudagraph() - def _can_run_prefill_cudagraph(self, infer_state, handle_token_num): + def _can_run_prefill_cudagraph(self, infer_state: DeepseekV4InferStateInfo, handle_token_num): if infer_state.prefix_total_token_num == 0: return True - if not self._logged_prefill_graph_prefix_skip: - logger.info( - "DeepSeek-V4 skips prefill cudagraph for prompt-cache extension batches; " - "no-prefix prefill batches still use prefill cudagraph." - ) - self._logged_prefill_graph_prefix_skip = True return False def _init_att_backend(self): From ffafdbf1ed8eec2530f8e541d5978fbc0161d952 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 8 Jun 2026 05:05:10 +0000 Subject: [PATCH 007/214] format --- lightllm/common/req_manager.py | 37 ++++++++-------------------------- 1 file changed, 8 insertions(+), 29 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index ca027a63c8..7b56129c3f 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -163,13 +163,11 @@ def __init__(self, max_request_num): ) elif self.penalty_counter_mode == "pin_mem_counter": self.req_to_out_token_id_counter = torch.zeros( - (max_request_num + 1, self.vocab_size), - dtype=torch.int32, - device="cpu", - pin_memory=True, + (max_request_num + 1, self.vocab_size), dtype=torch.int32, device="cpu", pin_memory=True ) 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()) self.req_to_presence_penalty[req.req_idx].fill_(shm_param.presence_penalty) @@ -199,18 +197,14 @@ def init_req_sampling_params(self, req: "InferReq"): dtype=torch.int32, ).cuda(non_blocking=True) token_id_counter( - prompt_ids=prompt_ids, - out_token_id_counter=self.req_to_out_token_id_counter[req.req_idx], + prompt_ids=prompt_ids, out_token_id_counter=self.req_to_out_token_id_counter[req.req_idx] ) torch.cuda.current_stream().synchronize() return def update_reqs_out_token_counter_gpu( - self, - b_req_idx: torch.Tensor, - next_token_ids: torch.Tensor, - mask: torch.Tensor = None, + self, b_req_idx: torch.Tensor, next_token_ids: torch.Tensor, mask: torch.Tensor = None ): if self.penalty_counter_mode not in ["gpu_counter", "pin_mem_counter"]: return @@ -226,10 +220,7 @@ def update_reqs_out_token_counter_gpu( return def update_reqs_token_counter( - self, - req_objs: List["InferReq"], - next_token_ids: List[int], - accept_mark: Optional[List[List[bool]]] = None, + self, req_objs: List["InferReq"], next_token_ids: List[int], accept_mark: Optional[List[List[bool]]] = None ): if self.penalty_counter_mode != "cpu_counter": return @@ -271,13 +262,7 @@ def gen_cpu_out_token_counter_sampling_params(self, req_objs: List["InferReq"]): class ReqManagerForMamba(ReqManager): - def __init__( - self, - max_request_num, - max_sequence_length, - mem_manager, - linear_config: LinearAttCacheConfig, - ): + def __init__(self, max_request_num, max_sequence_length, mem_manager, linear_config: LinearAttCacheConfig): super().__init__(max_request_num, max_sequence_length, mem_manager) self.mtp_step = get_env_start_args().mtp_step self.big_page_token_num = ( @@ -322,6 +307,7 @@ def get_mamba_cache(self, layer_idx_in_all: int): return conv_states, ssm_states def copy_big_page_buffer_to_linear_att_state(self, big_page_buffer_idx: int, req: "InferReq"): + from .linear_att_cache_manager import LinearAttCacheManager big_page_buffers: LinearAttCacheManager = self.mem_manager.linear_att_big_page_buffers @@ -375,15 +361,8 @@ def __init__( indexer_head_dim: Optional[int] = None, ): super().__init__(max_request_num, max_sequence_length, mem_manager) - if mem_manager is not None: - assert isinstance(mem_manager, DeepseekV4MemoryManager) - compress_rates = mem_manager.compress_rates - head_dim = mem_manager.head_dim - indexer_head_dim = mem_manager.indexer_head_dim - assert compress_rates is not None, "DeepSeek-V4 req manager requires compress_rates" - assert head_dim is not None, "DeepSeek-V4 req manager requires head_dim" - assert indexer_head_dim is not None, "DeepSeek-V4 req manager requires indexer_head_dim" + self.mem_manager = mem_manager self.compress_rates = list(compress_rates) self.n_c4 = sum(1 for r in self.compress_rates if r == 4) self.n_c128 = sum(1 for r in self.compress_rates if r == 128) From e8009cb3e053ffe7dbe465c027e4fe6a676181c8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 11 Jun 2026 02:15:12 +0000 Subject: [PATCH 008/214] pass gsm8k but need review --- .../attention/nsa/fp8_flashmla_sparse.py | 401 +++++- lightllm/common/basemodel/basemodel.py | 80 +- .../fused_moe/fused_moe_weight.py | 36 +- .../meta_weights/fused_moe/impl/__init__.py | 6 + .../meta_weights/fused_moe/impl/mxfp4_impl.py | 44 + .../fused_moe/impl/triton_impl.py | 19 + .../fused_moe/grouped_fused_moe_ep.py | 2 +- .../deepseek4_mem_manager.py | 1218 +++++++--------- lightllm/common/quantization/__init__.py | 9 +- lightllm/common/quantization/deepgemm.py | 78 + lightllm/common/req_manager.py | 693 ++++----- lightllm/models/deepseek_v4/infer_struct.py | 8 +- .../deepseek_v4/layer_infer/attention.py | 149 -- .../deepseek_v4/layer_infer/compressor.py | 645 ++++----- .../layer_infer/hyper_connection.py | 75 +- .../layer_infer/post_layer_infer.py | 8 +- .../layer_infer/pre_layer_infer.py | 19 +- .../layer_infer/transformer_layer_infer.py | 1251 ++++++----------- .../layer_weights/transformer_layer_weight.py | 331 +---- lightllm/models/deepseek_v4/mem_manager.py | 12 - lightllm/models/deepseek_v4/model.py | 86 +- .../destindex_copy_indexer_k_dsv4.py | 92 ++ .../destindex_copy_kv_flashmla_dsv4.py | 121 ++ .../triton_kernel/quant_convert.py | 77 - lightllm/server/api_cli.py | 7 +- .../router/dynamic_prompt/radix_cache.py | 57 + .../server/router/model_infer/infer_batch.py | 167 +-- .../model_infer/mode_backend/base_backend.py | 29 +- 28 files changed, 2591 insertions(+), 3129 deletions(-) create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py delete mode 100644 lightllm/models/deepseek_v4/layer_infer/attention.py delete mode 100644 lightllm/models/deepseek_v4/mem_manager.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 539ade769e..0570adea83 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -1,4 +1,5 @@ import dataclasses +import inspect import torch from typing import TYPE_CHECKING, Tuple @@ -9,6 +10,245 @@ from lightllm.common.basemodel.infer_struct import InferStateInfo +FLASHMLA_INDEX_ALIGN = 64 +# this flash_mla extra-cache fork only instantiates h_q in {64, 128}; pad TP-split q heads up +# to the nearest supported count (zero heads are discarded from the output slice). +FLASHMLA_SUPPORTED_HEADS = (64, 128) + + +def _pad_q_heads(q_4d: torch.Tensor, attn_sink: torch.Tensor): + h_q = q_4d.shape[2] + if h_q in FLASHMLA_SUPPORTED_HEADS: + return q_4d, attn_sink, h_q + target = next((h for h in FLASHMLA_SUPPORTED_HEADS if h >= h_q), None) + assert target is not None, f"num q heads {h_q} exceeds flash_mla support {FLASHMLA_SUPPORTED_HEADS}" + q_pad = torch.nn.functional.pad(q_4d, (0, 0, 0, target - h_q)) + sink_pad = torch.nn.functional.pad(attn_sink, (0, target - h_q)) + return q_pad, sink_pad, h_q + + +class DeepseekV4MissingOperatorError(RuntimeError): + pass + + +def _missing_attention_op(feature: str) -> None: + raise DeepseekV4MissingOperatorError( + f"DeepSeek-V4 {feature} has no production batch operator. The flashmla_kvcache path " + f"(packed swa/c4/c128 pools + paged compressor + indexer top-k) is the supported route; " + f"this legacy/non-flashmla entry point was never wired and is fenced on purpose." + ) + + +def _pad_last_dim(x: torch.Tensor, multiple: int = FLASHMLA_INDEX_ALIGN, value: int = -1) -> torch.Tensor: + pad = (-x.shape[-1]) % multiple + if pad == 0: + return x.contiguous() + out = torch.full((*x.shape[:-1], x.shape[-1] + pad), value, dtype=x.dtype, device=x.device) + out[..., : x.shape[-1]] = x + return out.contiguous() + + +def _view_dsv4_flashmla_cache(layer_buffer: torch.Tensor, page_size: int) -> torch.Tensor: + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_MLA_BYTES_PER_TOKEN + + usable = page_size * DSV4_MLA_BYTES_PER_TOKEN + return layer_buffer[:, :usable].view(layer_buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) + + +def _load_flash_mla_with_extra(): + try: + import flash_mla + except Exception as exc: + raise DeepseekV4MissingOperatorError( + "DeepSeek-V4 packed FlashMLA requires the flash_mla package with compiled CUDA extension. " + f"Import failed with: {type(exc).__name__}: {exc}" + ) from exc + + fn = getattr(flash_mla, "flash_mla_with_kvcache", None) + get_mla_metadata = getattr(flash_mla, "get_mla_metadata", None) + missing_symbols = [] + if fn is None: + missing_symbols.append("flash_mla_with_kvcache") + if get_mla_metadata is None: + missing_symbols.append("get_mla_metadata") + if missing_symbols: + raise DeepseekV4MissingOperatorError( + "DeepSeek-V4 requires flash_mla.flash_mla_with_kvcache extra-cache wrapper. " + f"Current module={getattr(flash_mla, '__file__', '')} " + f"is missing symbols {missing_symbols}." + ) + + sig = inspect.signature(fn) + required = { + "attn_sink", + "extra_k_cache", + "extra_indices_in_kvcache", + "topk_length", + "extra_topk_length", + } + missing = sorted(required.difference(sig.parameters)) + if missing: + raise DeepseekV4MissingOperatorError( + "DeepSeek-V4 requires flash_mla.flash_mla_with_kvcache with extra-cache arguments. " + f"Current module={getattr(flash_mla, '__file__', '')} is missing {missing}." + ) + return flash_mla + + +def _build_dsv4_repeated_prefill_reqs(infer_state) -> torch.Tensor: + return torch.repeat_interleave(infer_state.b_req_idx, infer_state.b_q_seq_len.long()) + + +def _build_dsv4_prefill_positions(infer_state) -> torch.Tensor: + total = infer_state.total_token_num - infer_state.prefix_total_token_num + token_offsets = torch.arange(total, dtype=torch.int32, device=infer_state.b_q_seq_len.device) + req_ids = torch.repeat_interleave( + torch.arange(infer_state.batch_size, dtype=torch.long, device=infer_state.b_q_seq_len.device), + infer_state.b_q_seq_len.long(), + ) + local_offsets = token_offsets - infer_state.b_q_start_loc[req_ids] + return infer_state.b_ready_cache_len[req_ids] + local_offsets + + +def _build_dsv4_swa_indices( + req_manager, + mem_manager, + req_idx: torch.Tensor, + positions: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + window = int(mem_manager.sliding_window) + offsets = positions[:, None] - torch.arange(window, dtype=positions.dtype, device=positions.device)[None, :] + valid_pos = offsets >= 0 + safe_offsets = offsets.clamp_min(0).long() + full_slots = req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_offsets] + swa_slots = mem_manager.full_to_swa_indexs[full_slots.long()].to(torch.int32) + indices = torch.where(valid_pos, swa_slots, torch.full_like(swa_slots, -1)) + lengths = torch.clamp(positions + 1, min=1, max=window).to(torch.int32) + return _pad_last_dim(indices.to(torch.int32)).unsqueeze(1), lengths.contiguous() + + +def _gather_dsv4_compress_slots( + infer_state, + mapping: torch.Tensor, + req_idx: torch.Tensor, + valid: torch.Tensor, + offsets: torch.Tensor, + ratio: int, +) -> torch.Tensor: + """条目 g 的压缩槽 = full_to_c*[req_to_token[req, (g+1)*ratio-1]](组末 token 的 full 槽位)。 + 无效条目(超出因果长度/HOLD 行)用位置 0 安全 gather 后由调用方按 valid 掩掉。""" + end_pos = offsets[None, :] * ratio + (ratio - 1) + safe_pos = torch.where(valid, end_pos, torch.zeros_like(end_pos)) + full_slots = infer_state.req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_pos] + return mapping[full_slots.long()].to(torch.int32) + + +def _build_dsv4_c128_indices( + infer_state, + req_idx: torch.Tensor, + positions: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + raw_lengths = (positions + 1) // 128 + lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) + max_len = max(1, int(infer_state.max_kv_seq_len) // 128) + offsets = torch.arange(max_len, dtype=torch.long, device=positions.device) + valid = offsets[None, :] < raw_lengths[:, None] + slots = _gather_dsv4_compress_slots( + infer_state, infer_state.mem_manager.full_to_c128_indexs, req_idx, valid, offsets, 128 + ) + indices = torch.where(valid, slots, torch.full_like(slots, -1)) + return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() + + +def _build_dsv4_c4_indices( + infer_state, + layer_index: int, + req_idx: torch.Tensor, + positions: torch.Tensor, + nsa_dict: dict, +) -> Tuple[torch.Tensor, torch.Tensor]: + """c4(CSA) extra indices: causal all-entries when the entry space fits index_topk, + otherwise Lightning-Indexer scored top-k. Pure tensor ops (decode runs inside cuda graphs).""" + import torch.distributed as dist + import torch.nn.functional as F + from lightllm.distributed.communication_op import all_reduce + + mem_manager = infer_state.mem_manager + raw_lengths = (positions + 1) // 4 + max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) + index_topk = int(nsa_dict["index_topk"]) + offsets = torch.arange(max_entries, dtype=torch.long, device=positions.device) + valid = offsets[None, :] < raw_lengths[:, None] + slots = _gather_dsv4_compress_slots(infer_state, mem_manager.full_to_c4_indexs, req_idx, valid, offsets, 4) + + if max_entries <= index_topk: + lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) + indices = torch.where(valid, slots, torch.full_like(slots, -1)) + return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() + + idx_q = nsa_dict["idx_q"] # [T, H, index_head_dim], rope applied + idx_weight = nsa_dict["idx_weight"] # [T, H] fp32, weight scale applied + score_scale = float(nsa_dict["indexer_score_scale"]) + hold_slot = mem_manager.c4_indexer_pool.HOLD_TOKEN_MEMINDEX + safe_slots = torch.where(valid, slots.long(), torch.full_like(slots.long(), hold_slot)) + k = mem_manager.gather_indexer_k(layer_index, safe_slots.reshape(-1)).view(positions.shape[0], max_entries, -1) + + num_tokens, num_heads = idx_q.shape[0], idx_q.shape[1] + score_chunks = [] + chunk = max(1, min(num_tokens, (16 * 1024 * 1024) // max(1, num_heads * max_entries))) + for start in range(0, num_tokens, chunk): + end = min(num_tokens, start + chunk) + scores = torch.einsum("thd,tnd->thn", idx_q[start:end].float(), k[start:end].float()) + scores = F.relu(scores) * score_scale + score_chunks.append((scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1)) + index_scores = torch.cat(score_chunks, dim=0) + if int(nsa_dict.get("tp_world_size", 1)) > 1: + all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + index_scores = index_scores.masked_fill(~valid, float("-inf")) + top = index_scores.topk(index_topk, dim=-1).indices + top_valid = torch.gather(valid, 1, top) + top_slots = torch.gather(slots.long(), 1, top).to(torch.int32) + indices = torch.where(top_valid, top_slots, torch.full_like(top_slots, -1)) + lengths = torch.clamp(torch.minimum(raw_lengths, torch.full_like(raw_lengths, index_topk)), min=1) + return _pad_last_dim(indices).unsqueeze(1), lengths.to(torch.int32).contiguous() + + +def _build_dsv4_extra_metadata( + infer_state, + layer_index: int, + compress_ratio: int, + req_idx: torch.Tensor, + positions: torch.Tensor, + swa_indices: torch.Tensor, + swa_lengths: torch.Tensor, + nsa_dict: dict, +) -> "_Dsv4Metadata": + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C128_PAGE_SIZE, DSV4_C4_PAGE_SIZE + + if compress_ratio == 0: + return _Dsv4Metadata(swa_indices, swa_lengths) + if compress_ratio == 4: + extra_indices, extra_lengths = _build_dsv4_c4_indices(infer_state, layer_index, req_idx, positions, nsa_dict) + extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(layer_index) + extra_cache = _view_dsv4_flashmla_cache(extra_buffer, DSV4_C4_PAGE_SIZE) + return _Dsv4Metadata(swa_indices, swa_lengths, extra_cache, extra_indices, extra_lengths) + if compress_ratio == 128: + extra_indices, extra_lengths = _build_dsv4_c128_indices(infer_state, req_idx, positions) + extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(layer_index) + extra_cache = _view_dsv4_flashmla_cache(extra_buffer, DSV4_C128_PAGE_SIZE) + return _Dsv4Metadata(swa_indices, swa_lengths, extra_cache, extra_indices, extra_lengths) + raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") + + +@dataclasses.dataclass +class _Dsv4Metadata: + swa_indices: torch.Tensor + swa_lengths: torch.Tensor + extra_cache: torch.Tensor = None + extra_indices: torch.Tensor = None + extra_lengths: torch.Tensor = None + + class NsaFlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) @@ -17,6 +257,7 @@ def __init__(self, model): torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) for _ in range(2) ] + self._flash_mla = None def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparsePrefillAttState": return NsaFlashMlaFp8SparsePrefillAttState(backend=self, infer_state=infer_state) @@ -24,6 +265,11 @@ def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMl def create_att_decode_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparseDecodeAttState": return NsaFlashMlaFp8SparseDecodeAttState(backend=self, infer_state=infer_state) + def flash_mla(self): + if self._flash_mla is None: + self._flash_mla = _load_flash_mla_with_extra() + return self._flash_mla + @dataclasses.dataclass class NsaFlashMlaFp8SparsePrefillAttState(BasePrefillAttState): @@ -62,6 +308,12 @@ def prefill_att( ) -> torch.Tensor: assert att_control.nsa_prefill, "nsa_prefill must be True for NSA prefill attention" assert att_control.nsa_prefill_dict is not None, "nsa_prefill_dict is required" + if att_control.nsa_prefill_dict.get("flashmla_kvcache"): + return self._flashmla_kvcache_prefill_att( + q=q, + packed_kv=k, + nsa_dict=att_control.nsa_prefill_dict, + ) return self._nsa_prefill_att(q=q, packed_kv=k, att_control=att_control) def _nsa_prefill_att( @@ -78,6 +330,8 @@ def _nsa_prefill_att( kv_lora_rank = nsa_dict["kv_lora_rank"] topk_mem_indices = nsa_dict["topk_mem_indices"] prefill_cache_kv = nsa_dict["prefill_cache_kv"] + attn_sink = nsa_dict.get("attn_sink") + topk_length = nsa_dict.get("topk_length") if self.infer_state.prefix_total_token_num > 0: # 当前推理生成的token kv部分从 prefill_cache_kv 中获取,历史 @@ -101,9 +355,72 @@ def _nsa_prefill_att( indices=topk_indices, sm_scale=softmax_scale, d_v=kv_lora_rank, + attn_sink=attn_sink, + topk_length=topk_length, ) return mla_out + def _build_flashmla_kvcache_prefill_metadata(self, nsa_dict: dict) -> _Dsv4Metadata: + infer_state = self.infer_state + req_idx = _build_dsv4_repeated_prefill_reqs(infer_state) + positions = _build_dsv4_prefill_positions(infer_state) + swa_indices, swa_lengths = _build_dsv4_swa_indices( + infer_state.req_manager, + infer_state.mem_manager, + req_idx, + positions, + ) + return _build_dsv4_extra_metadata( + infer_state, + nsa_dict["layer_index"], + nsa_dict["compress_ratio"], + req_idx, + positions, + swa_indices, + swa_lengths, + nsa_dict, + ) + + def _flashmla_kvcache_prefill_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: + attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() + metadata = self._build_flashmla_kvcache_prefill_metadata(nsa_dict) + return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) + + def _flashmla_kvcache_att( + self, + q: torch.Tensor, + packed_kv: torch.Tensor, + metadata: _Dsv4Metadata, + attn_sink: torch.Tensor, + nsa_dict: dict, + ) -> torch.Tensor: + flash_mla = self.backend.flash_mla() + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE + + q_4d = q.unsqueeze(1).contiguous() + q_4d, attn_sink, num_real_heads = _pad_q_heads(q_4d, attn_sink) + k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) + sched_meta, _ = flash_mla.get_mla_metadata() + out, _ = flash_mla.flash_mla_with_kvcache( + q=q_4d, + k_cache=k_cache, + block_table=None, + cache_seqlens=None, + head_dim_v=nsa_dict["head_dim_v"], + tile_scheduler_metadata=sched_meta, + num_splits=None, + softmax_scale=nsa_dict["softmax_scale"], + causal=False, + is_fp8_kvcache=True, + indices=metadata.swa_indices, + attn_sink=attn_sink, + topk_length=metadata.swa_lengths, + extra_k_cache=metadata.extra_cache, + extra_indices_in_kvcache=metadata.extra_indices, + extra_topk_length=metadata.extra_lengths, + ) + return out[:, 0, :num_real_heads].contiguous() + @dataclasses.dataclass class NsaFlashMlaFp8SparseDecodeAttState(BaseDecodeAttState): @@ -141,9 +458,10 @@ def init_state(self): ragged_mem_index=self.ragged_mem_index, hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, ) - import flash_mla - - self.flashmla_sched_meta, _ = flash_mla.get_mla_metadata() + flash_mla = self.backend.flash_mla() + # one sched_meta per layer type: the lazy config locks extra-cache geometry (page size, + # presence) on first invocation, so swa-only/c4/c128 layers must not share one object. + self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} return def decode_att( @@ -156,6 +474,12 @@ def decode_att( ) -> torch.Tensor: assert att_control.nsa_decode, "nsa_decode must be True for NSA decode attention" assert att_control.nsa_decode_dict is not None, "nsa_decode_dict is required" + if att_control.nsa_decode_dict.get("flashmla_kvcache"): + return self._flashmla_kvcache_decode_att( + q=q, + packed_kv=k, + nsa_dict=att_control.nsa_decode_dict, + ) return self._nsa_decode_att(q=q, packed_kv=k, att_control=att_control) def _nsa_decode_att( @@ -170,6 +494,11 @@ def _nsa_decode_att( topk_mem_indices = nsa_dict["topk_mem_indices"] softmax_scale = nsa_dict["softmax_scale"] kv_lora_rank = nsa_dict["kv_lora_rank"] + attn_sink = nsa_dict.get("attn_sink") + topk_length = nsa_dict.get("topk_length") + extra_k_cache = nsa_dict.get("extra_k_cache") + extra_indices = nsa_dict.get("extra_indices_in_kvcache") + extra_topk_length = nsa_dict.get("extra_topk_length") if topk_mem_indices.ndim == 2: topk_mem_indices = topk_mem_indices.unsqueeze(1) @@ -189,10 +518,74 @@ def _nsa_decode_att( block_table=None, cache_seqlens=None, head_dim_v=kv_lora_rank, - tile_scheduler_metadata=self.flashmla_sched_meta, + tile_scheduler_metadata=self.flashmla_sched_meta[0], softmax_scale=softmax_scale, causal=False, is_fp8_kvcache=True, indices=topk_mem_indices, + attn_sink=attn_sink, + topk_length=topk_length, + extra_k_cache=extra_k_cache, + extra_indices_in_kvcache=extra_indices, + extra_topk_length=extra_topk_length, ) return o_tensor[:, 0, :, :] # [b, 1, h, d] -> [b, h, d] + + def _build_flashmla_kvcache_decode_metadata(self, nsa_dict: dict) -> _Dsv4Metadata: + infer_state = self.infer_state + positions = infer_state.b_seq_len.to(torch.int32) - 1 + swa_indices, swa_lengths = _build_dsv4_swa_indices( + infer_state.req_manager, + infer_state.mem_manager, + infer_state.b_req_idx, + positions, + ) + return _build_dsv4_extra_metadata( + infer_state, + nsa_dict["layer_index"], + nsa_dict["compress_ratio"], + infer_state.b_req_idx, + positions, + swa_indices, + swa_lengths, + nsa_dict, + ) + + def _flashmla_kvcache_decode_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: + attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() + metadata = self._build_flashmla_kvcache_decode_metadata(nsa_dict) + return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) + + def _flashmla_kvcache_att( + self, + q: torch.Tensor, + packed_kv: torch.Tensor, + metadata: _Dsv4Metadata, + attn_sink: torch.Tensor, + nsa_dict: dict, + ) -> torch.Tensor: + flash_mla = self.backend.flash_mla() + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE + + q_4d = q.unsqueeze(1).contiguous() + q_4d, attn_sink, num_real_heads = _pad_q_heads(q_4d, attn_sink) + k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) + out, _ = flash_mla.flash_mla_with_kvcache( + q=q_4d, + k_cache=k_cache, + block_table=None, + cache_seqlens=None, + head_dim_v=nsa_dict["head_dim_v"], + tile_scheduler_metadata=self.flashmla_sched_meta[nsa_dict["compress_ratio"]], + num_splits=None, + softmax_scale=nsa_dict["softmax_scale"], + causal=False, + is_fp8_kvcache=True, + indices=metadata.swa_indices, + attn_sink=attn_sink, + topk_length=metadata.swa_lengths, + extra_k_cache=metadata.extra_cache, + extra_indices_in_kvcache=metadata.extra_indices, + extra_topk_length=metadata.extra_lengths, + ) + return out[:, 0, :num_real_heads].contiguous() diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 8e352519c0..986802e760 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -527,13 +527,17 @@ def _prefill( alloc_mem_index=infer_state.mem_index, max_q_seq_len=infer_state.max_q_seq_len, ) - if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): - self.mem_manager.prepare_prefill_swa_slots( + if hasattr(self.req_manager, "prepare_prefill_swa"): + self.req_manager.prepare_prefill_swa( b_req_idx=infer_state.b_req_idx, + b_ready_cache_len=infer_state.b_ready_cache_len, b_seq_len=infer_state.b_seq_len, + ) + if hasattr(self.req_manager, "prepare_prefill_compress_slots"): + self.req_manager.prepare_prefill_compress_slots( + b_req_idx=infer_state.b_req_idx, b_ready_cache_len=infer_state.b_ready_cache_len, - b_start_loc=model_input.b_prefill_start_loc, - mem_index=infer_state.mem_index, + b_seq_len=infer_state.b_seq_len, ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -541,8 +545,6 @@ def _prefill( infer_state.init_some_extra_state(self) infer_state.init_att_state() model_output = self._context_forward(infer_state) - if hasattr(self.mem_manager, "commit_prefill_swa_slots"): - self.mem_manager.commit_prefill_swa_slots() model_output = self._create_unpad_prefill_model_output( padded_model_output=model_output, @@ -577,12 +579,14 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) - if hasattr(self.mem_manager, "prepare_decode_swa_slots"): - self.mem_manager.prepare_decode_swa_slots( + if hasattr(self.req_manager, "prepare_decode_swa"): + self.req_manager.prepare_decode_swa( model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes ) if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots(model_input.b_req_idx, model_input.b_seq_len) + self.req_manager.prepare_decode_compress_slots( + model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes + ) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -604,12 +608,14 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) - if hasattr(self.mem_manager, "prepare_decode_swa_slots"): - self.mem_manager.prepare_decode_swa_slots( + if hasattr(self.req_manager, "prepare_decode_swa"): + self.req_manager.prepare_decode_swa( model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes ) if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots(model_input.b_req_idx, model_input.b_seq_len) + self.req_manager.prepare_decode_compress_slots( + model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes + ) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -775,13 +781,17 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state0.mem_index, max_q_seq_len=infer_state0.max_q_seq_len, ) - if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): - self.mem_manager.prepare_prefill_swa_slots( + if hasattr(self.req_manager, "prepare_prefill_swa"): + self.req_manager.prepare_prefill_swa( b_req_idx=infer_state0.b_req_idx, + b_ready_cache_len=infer_state0.b_ready_cache_len, b_seq_len=infer_state0.b_seq_len, + ) + if hasattr(self.req_manager, "prepare_prefill_compress_slots"): + self.req_manager.prepare_prefill_compress_slots( + b_req_idx=infer_state0.b_req_idx, b_ready_cache_len=infer_state0.b_ready_cache_len, - b_start_loc=model_input0.b_prefill_start_loc, - mem_index=infer_state0.mem_index, + b_seq_len=infer_state0.b_seq_len, ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -796,13 +806,17 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state1.mem_index, max_q_seq_len=infer_state1.max_q_seq_len, ) - if hasattr(self.mem_manager, "prepare_prefill_swa_slots"): - self.mem_manager.prepare_prefill_swa_slots( + if hasattr(self.req_manager, "prepare_prefill_swa"): + self.req_manager.prepare_prefill_swa( b_req_idx=infer_state1.b_req_idx, + b_ready_cache_len=infer_state1.b_ready_cache_len, b_seq_len=infer_state1.b_seq_len, + ) + if hasattr(self.req_manager, "prepare_prefill_compress_slots"): + self.req_manager.prepare_prefill_compress_slots( + b_req_idx=infer_state1.b_req_idx, b_ready_cache_len=infer_state1.b_ready_cache_len, - b_start_loc=model_input1.b_prefill_start_loc, - mem_index=infer_state1.mem_index, + b_seq_len=infer_state1.b_seq_len, ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -811,8 +825,6 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod prefill_mem_indexes_ready_event.record() model_output0, model_output1 = self._overlap_tpsp_context_forward(infer_state0, infer_state1=infer_state1) - if hasattr(self.mem_manager, "commit_prefill_swa_slots"): - self.mem_manager.commit_prefill_swa_slots() model_output0 = self._create_unpad_prefill_model_output( padded_model_output=model_output0, @@ -864,19 +876,19 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode # 一致,需要按照较高 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) - if hasattr(self.mem_manager, "prepare_decode_swa_slots"): - self.mem_manager.prepare_decode_swa_slots( + if hasattr(self.req_manager, "prepare_decode_swa"): + self.req_manager.prepare_decode_swa( padded_model_input0.b_req_idx, padded_model_input0.b_seq_len, padded_model_input0.mem_indexes ) - self.mem_manager.prepare_decode_swa_slots( + self.req_manager.prepare_decode_swa( padded_model_input1.b_req_idx, padded_model_input1.b_seq_len, padded_model_input1.mem_indexes ) if hasattr(self.req_manager, "prepare_decode_compress_slots"): self.req_manager.prepare_decode_compress_slots( - padded_model_input0.b_req_idx, padded_model_input0.b_seq_len + padded_model_input0.b_req_idx, padded_model_input0.b_seq_len, padded_model_input0.mem_indexes ) self.req_manager.prepare_decode_compress_slots( - padded_model_input1.b_req_idx, padded_model_input1.b_seq_len + padded_model_input1.b_req_idx, padded_model_input1.b_seq_len, padded_model_input1.mem_indexes ) infer_state0 = self._create_inferstate(padded_model_input0, 0) copy_kv_index_to_req( @@ -919,16 +931,20 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode 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) - if hasattr(self.mem_manager, "prepare_decode_swa_slots"): - self.mem_manager.prepare_decode_swa_slots( + if hasattr(self.req_manager, "prepare_decode_swa"): + self.req_manager.prepare_decode_swa( model_input0.b_req_idx, model_input0.b_seq_len, model_input0.mem_indexes ) - self.mem_manager.prepare_decode_swa_slots( + self.req_manager.prepare_decode_swa( model_input1.b_req_idx, model_input1.b_seq_len, model_input1.mem_indexes ) if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots(model_input0.b_req_idx, model_input0.b_seq_len) - self.req_manager.prepare_decode_compress_slots(model_input1.b_req_idx, model_input1.b_seq_len) + self.req_manager.prepare_decode_compress_slots( + model_input0.b_req_idx, model_input0.b_seq_len, model_input0.mem_indexes + ) + self.req_manager.prepare_decode_compress_slots( + model_input1.b_req_idx, model_input1.b_seq_len, model_input1.mem_indexes + ) infer_state0 = self._create_inferstate(model_input0, 0) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index fca9b80fcf..24842ed383 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -68,12 +68,13 @@ def __init__( auto_update_redundancy_expert=self.auto_update_redundancy_expert, ) self.lock = threading.Lock() + self._moe_weight_finalized = False self._create_weight() def _init_config(self, network_config: Dict[str, Any]): self.n_group = network_config.get("n_group", 0) self.use_grouped_topk = self.n_group > 0 - self.norm_topk_prob = network_config["norm_topk_prob"] + self.norm_topk_prob = network_config.get("norm_topk_prob", False) self.topk_group = network_config.get("topk_group", 0) self.num_experts_per_tok = network_config["num_experts_per_tok"] self.routed_scaling_factor = network_config.get("routed_scaling_factor", 1.0) @@ -136,6 +137,7 @@ def experts( is_prefill: Optional[bool] = None, ) -> torch.Tensor: """Backward compatible method that routes to platform-specific implementation.""" + self._finalize_moe_weight() return self.fuse_moe_impl( input_tensor=input_tensor, router_logits=router_logits, @@ -152,6 +154,25 @@ def experts( per_expert_scale=self.per_expert_scale, ) + def experts_with_preselected( + self, + input_tensor: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + ) -> torch.Tensor: + self._finalize_moe_weight() + return self.fuse_moe_impl.fused_experts_with_topk( + input_tensor=input_tensor, + w13=self.w13, + w2=self.w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + is_prefill=is_prefill, + clamp_limit=clamp_limit, + ) + def low_latency_dispatch( self, hidden_states: torch.Tensor, @@ -280,7 +301,18 @@ def verify_load(self): e_score_correction_bias_load_ok = ( True if self.e_score_correction_bias is None else getattr(self.e_score_correction_bias, "load_ok", False) ) - return weight_load_ok and per_expert_scale_load_ok and e_score_correction_bias_load_ok + load_ok = weight_load_ok and per_expert_scale_load_ok and e_score_correction_bias_load_ok + if load_ok: + self._finalize_moe_weight() + return load_ok + + def _finalize_moe_weight(self): + if self._moe_weight_finalized: + return + finalize = getattr(self.quant_method, "finalize_moe_weight", None) + if finalize is not None: + finalize(self) + self._moe_weight_finalized = True def _create_weight(self): intermediate_size = self.split_inter_size diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 67bb90e4ef..282c0abdce 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -2,9 +2,15 @@ from .triton_impl import FuseMoeTriton from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM +from .mxfp4_impl import FuseMoeMXFP4 def select_fuse_moe_impl(quant_method: QuantizationMethod, enable_ep_moe: bool): + if quant_method.method_name == "marlin-mxfp4w4a16-b32": + if enable_ep_moe: + raise RuntimeError("marlin-mxfp4w4a16-b32 does not support enable_ep_moe yet") + return FuseMoeMXFP4 + if enable_ep_moe: return FuseMoeDeepGEMM diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py new file mode 100644 index 0000000000..97cf238116 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py @@ -0,0 +1,44 @@ +import torch +from typing import Optional + +from lightllm.common.quantization.quantize_method import WeightPack +from .triton_impl import FuseMoeTriton + + +class FuseMoeMXFP4(FuseMoeTriton): + def create_workspace(self): + return None + + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + ): + try: + from vllm.model_executor.layers.fused_moe.activation import MoEActivation + from vllm.model_executor.layers.fused_moe.experts.marlin_moe import fused_marlin_moe + from vllm.scalar_type import scalar_types + except Exception as e: + raise RuntimeError(f"MXFP4 fused MoE requires vLLM fused kernels, error={repr(e)}") from e + + return fused_marlin_moe( + hidden_states=input_tensor.contiguous(), + w1=w13.weight, + w2=w2.weight, + bias1=None, + bias2=None, + w1_scale=w13.weight_scale, + w2_scale=w2.weight_scale, + topk_weights=topk_weights.to(torch.float32).contiguous(), + topk_ids=topk_ids.to(torch.long).contiguous(), + quant_type_id=scalar_types.float4_e2m1f.id, + global_num_experts=self.n_routed_experts, + activation=MoEActivation.SILU, + clamp_limit=clamp_limit, + ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index a0d30547a3..09ce88e3fd 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -114,6 +114,25 @@ def _fused_experts( ) return input_tensor + def fused_experts_with_topk( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + ): + return self._fused_experts( + input_tensor=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + is_prefill=is_prefill, + ) + def __call__( self, input_tensor: torch.Tensor, diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index cb2e370cb9..28fe6e4304 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -49,7 +49,7 @@ def check_ep_expert_dtype(quant_method: Any): "EP MoE requires --expert_dtype to be one of ['fp8', 'fp4'], " f"but the resolved fused_moe quant method is `{expert_dtype}`. " "Please start with --expert_dtype fp8 or --expert_dtype fp4. " - "Note that --expert_dtype fp4 is only supported on SM100 GPUs." + "Note that --expert_dtype fp4 with EP MoE is only supported on SM100 GPUs." ) if expert_dtype == "deepgemm-fp4fp8-b32" and not is_sm100_gpu(): raise RuntimeError( diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index dc708e0790..47fdf76fd9 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,7 +1,7 @@ import torch import torch.distributed as dist -from typing import Dict, List, Optional -from .deepseek2_mem_manager import Deepseek2MemoryManager +from typing import List, Optional, Union +from .mem_manager import MemoryManager from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator from lightllm.utils.dist_utils import get_current_device_id, get_current_rank_in_node @@ -12,36 +12,46 @@ logger = init_logger(__name__) +# fp8_ds_mla packed-latent byte layout (ABI shared with the flash_mla extra-cache fork and +# sglang/vllm): 448B NoPE fp8 + 64*2B RoPE bf16 + 7B ue8m0 scale + 1B pad = 584B per token, +# stored in page slabs whose tail carries the per-token scale bytes. DSV4_MLA_NOPE_DIM = 448 DSV4_MLA_ROPE_DIM = 64 DSV4_MLA_HEAD_DIM = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM DSV4_MLA_QUANT_GROUP_SIZE = 64 DSV4_MLA_SCALE_BYTES = DSV4_MLA_NOPE_DIM // DSV4_MLA_QUANT_GROUP_SIZE + 1 DSV4_MLA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 + DSV4_MLA_SCALE_BYTES +DSV4_MLA_DATA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 +DSV4_MLA_PAGE_ALIGN_BYTES = DSV4_MLA_DATA_BYTES_PER_TOKEN DSV4_INDEXER_HEAD_DIM = 128 -DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + 4 +DSV4_INDEXER_SCALE_BYTES = 4 +DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + DSV4_INDEXER_SCALE_BYTES DSV4_FP8_E4M3_MAX = 448.0 DSV4_FP8_SCALE_MIN = 1e-4 -DSV4_MLA_DATA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 -DSV4_MLA_SCALE_TAIL_BYTES = DSV4_MLA_SCALE_BYTES -DSV4_MLA_PAGE_ALIGN_BYTES = DSV4_MLA_DATA_BYTES_PER_TOKEN DSV4_SWA_PAGE_SIZE = 128 DSV4_C4_PAGE_SIZE = 64 DSV4_C128_PAGE_SIZE = 2 +# c4 compressor state ring(overlap 对: 每页 2 个分组槽 × ratio 4 行)。c128 state 在 128 边界 +# 自然归零(在线聚合),无缓存常驻需求,保持 req 键控,不进 swa 派生池。 +DSV4_C4_STATE_RING = 8 DSV4_PROFILE_MAX_FULL_TOKENS = 1_500_000 +# swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 的对应物)。 +# lightllm 的调度准入只看 full 池,prefill 优先的波次会让"已 prefill 未 decode"的请求整段 +# prompt 占住 swa 槽(首次 decode prep 才批量出窗回收),峰值≈准入波次 prompt 总和。在 +# v5 的 swa 压力阀/准入耦合落地前,用比 sglang 更宽的 0.3 兜住该瞬时峰值。 +DSV4_SWA_FULL_TOKENS_RATIO = 0.3 def _ceil_div(a: int, b: int) -> int: return (a + b - 1) // b -class _PageSlabMlaPool: - """SGLang-compatible fp8_ds_mla page-slab storage with token-slot addressing. +class PackedPagePool: + """fp8_ds_mla 风格的 page-slab 存储: 每页前段连续放 token 的 data 字节,页尾放 per-token scale 字节。 - The public loc is still a LightLLM token slot. Internally each page stores all - 576B NoPE+RoPE payloads first and the 8B scale records at the page tail: - data_offset = page * bytes_per_page + token_in_page * 576 - scale_offset = page * bytes_per_page + page_size * 576 + token_in_page * 8 + 寻址是纯 token 槽位 (page = slot // page_size),page 只是 scale-tail/对齐的物理打包技巧, + 不存在页粒度的分配。``write``/``read`` 是 torch 参考实现(单测 oracle);生产写入走 + triton packed writer(destindex_copy_kv_flashmla_dsv4 等),kernel 直接消费 ``buffer``。 """ def __init__( @@ -49,27 +59,26 @@ def __init__( size: int, page_size: int, layer_num: int, + data_bytes: int, + scale_bytes: int, + align_bytes: int = 1, device: str = "cuda", ): self.size = size self.page_size = page_size self.layer_num = layer_num - self.dtype = torch.uint8 - self.data_bytes_per_token = DSV4_MLA_DATA_BYTES_PER_TOKEN - self.scale_bytes_per_token = DSV4_MLA_SCALE_TAIL_BYTES - self.bytes_per_token = DSV4_MLA_BYTES_PER_TOKEN + self.data_bytes_per_token = data_bytes + self.scale_bytes_per_token = scale_bytes + self.bytes_per_token = data_bytes + scale_bytes self.num_pages = _ceil_div(size + 1, page_size) - self.bytes_per_page = ( - _ceil_div(page_size * self.bytes_per_token, DSV4_MLA_PAGE_ALIGN_BYTES) * DSV4_MLA_PAGE_ALIGN_BYTES - ) - self.scale_offset_in_page = page_size * self.data_bytes_per_token - self.kv_buffer = torch.zeros( - (layer_num, self.num_pages, self.bytes_per_page), - dtype=torch.uint8, - device=device, - ) + self.bytes_per_page = _ceil_div(page_size * self.bytes_per_token, align_bytes) * align_bytes + self.scale_offset_in_page = page_size * data_bytes + self.buffer = torch.zeros((layer_num, self.num_pages, self.bytes_per_page), dtype=torch.uint8, device=device) self.HOLD_TOKEN_MEMINDEX = size + def get_layer_buffer(self, layer_index: int) -> torch.Tensor: + return self.buffer[layer_index] + def _loc_offsets(self, loc: torch.Tensor): loc = loc.long() page = torch.div(loc, self.page_size, rounding_mode="floor") @@ -82,24 +91,21 @@ def _loc_offsets(self, loc: torch.Tensor): def write(self, layer_index: int, loc: torch.Tensor, packed: torch.Tensor) -> None: if loc.numel() == 0: return - loc = loc.long() - packed = packed.reshape(-1, DSV4_MLA_BYTES_PER_TOKEN).contiguous() - flat = self.kv_buffer[layer_index].view(-1) + loc = loc.reshape(-1) + packed = packed.reshape(-1, self.bytes_per_token).contiguous() + flat = self.buffer[layer_index].view(-1) data_offsets, scale_offsets = self._loc_offsets(loc) - - data = packed[:, : self.data_bytes_per_token].contiguous() - scale = packed[:, self.data_bytes_per_token : self.bytes_per_token].contiguous() data_range = torch.arange(self.data_bytes_per_token, device=loc.device) scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device) - flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)] = data - flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] = scale + flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)] = packed[:, : self.data_bytes_per_token] + flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] = packed[:, self.data_bytes_per_token :] return def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: - loc = loc.long() + loc = loc.reshape(-1) if loc.numel() == 0: - return torch.empty((0, DSV4_MLA_BYTES_PER_TOKEN), dtype=torch.uint8, device=self.kv_buffer.device) - flat = self.kv_buffer[layer_index].view(-1) + return torch.empty((0, self.bytes_per_token), dtype=torch.uint8, device=self.buffer.device) + flat = self.buffer[layer_index].view(-1) data_offsets, scale_offsets = self._loc_offsets(loc) data_range = torch.arange(self.data_bytes_per_token, device=loc.device) scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device) @@ -107,150 +113,24 @@ def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] return torch.cat([data, scale], dim=1).contiguous() - def get_layer_buffer(self, layer_index: int) -> torch.Tensor: - return self.kv_buffer[layer_index] - - -class _PageSlabIndexerPool: - """C4 indexer-K storage: page tail stores per-token fp32 scales.""" - - def __init__( - self, - size: int, - page_size: int, - layer_num: int, - device: str = "cuda", - ): - self.size = size - self.page_size = page_size - self.layer_num = layer_num - self.head_dim = DSV4_INDEXER_HEAD_DIM - self.scale_bytes = 4 - self.bytes_per_token = DSV4_INDEXER_BYTES_PER_TOKEN - self.num_pages = _ceil_div(size + 1, page_size) - self.bytes_per_page = page_size * self.bytes_per_token - self.scale_offset_in_page = page_size * self.head_dim - self.index_k_buffer = torch.zeros( - (layer_num, self.num_pages, self.bytes_per_page), - dtype=torch.uint8, - device=device, - ) - self.HOLD_TOKEN_MEMINDEX = size - - def _loc_offsets(self, loc: torch.Tensor): - loc = loc.long() - page = torch.div(loc, self.page_size, rounding_mode="floor") - token = loc % self.page_size - page_base = page * self.bytes_per_page - k_offsets = page_base + token * self.head_dim - scale_offsets = page_base + self.scale_offset_in_page + token * self.scale_bytes - return k_offsets, scale_offsets - - def write(self, layer_index: int, loc: torch.Tensor, packed: torch.Tensor) -> None: - if loc.numel() == 0: - return - loc = loc.long() - packed = packed.reshape(-1, self.bytes_per_token).contiguous() - flat = self.index_k_buffer[layer_index].view(-1) - k_offsets, scale_offsets = self._loc_offsets(loc) - k_range = torch.arange(self.head_dim, device=loc.device) - scale_range = torch.arange(self.scale_bytes, device=loc.device) - flat[k_offsets.unsqueeze(1) + k_range.unsqueeze(0)] = packed[:, : self.head_dim] - flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] = packed[:, self.head_dim :] - return - - def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: - loc = loc.long() - if loc.numel() == 0: - return torch.empty((0, self.bytes_per_token), dtype=torch.uint8, device=self.index_k_buffer.device) - flat = self.index_k_buffer[layer_index].view(-1) - k_offsets, scale_offsets = self._loc_offsets(loc) - k_range = torch.arange(self.head_dim, device=loc.device) - scale_range = torch.arange(self.scale_bytes, device=loc.device) - k = flat[k_offsets.unsqueeze(1) + k_range.unsqueeze(0)] - scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] - return torch.cat([k, scale], dim=1).contiguous() - - def get_layer_buffer(self, layer_index: int) -> torch.Tensor: - return self.index_k_buffer[layer_index] - - -class _SubKvPool: - """Compressed c4/c128 KV pool with token-slot allocator and page-slab backing.""" - - def __init__( - self, - size: int, - page_size: int, - layer_num: int, - with_indexer: bool = False, - shared_name: Optional[str] = None, - device: str = "cuda", - ): - self.size = size - self.dtype = torch.uint8 - self.layer_num = layer_num - self.page_size = page_size - self.mla_pool = _PageSlabMlaPool(size=size, page_size=page_size, layer_num=layer_num, device=device) - self.kv_buffer = self.mla_pool.kv_buffer - if with_indexer: - self.indexer_pool = _PageSlabIndexerPool( - size=size, - page_size=page_size, - layer_num=layer_num, - device=device, - ) - self.index_k_buffer = self.indexer_pool.index_k_buffer - else: - self.indexer_pool = None - self.index_k_buffer = None - - self.allocator = KvCacheAllocator(size, shared_name=shared_name) - self.HOLD_TOKEN_MEMINDEX = size - - def alloc(self, need_size) -> torch.Tensor: - return self.allocator.alloc(need_size) - - def free(self, free_index) -> None: - self.allocator.free(free_index) - - def free_all(self) -> None: - self.allocator.free_all() - - def get_kv_buffer(self, layer_index: int) -> torch.Tensor: - return self.mla_pool.get_layer_buffer(layer_index) - - def get_index_k_buffer(self, layer_index: int) -> torch.Tensor: - assert self.indexer_pool is not None, "this sub pool has no indexer-K buffer" - return self.indexer_pool.get_layer_buffer(layer_index) - - def write_kv(self, layer_index: int, slots: torch.Tensor, packed: torch.Tensor) -> None: - self.mla_pool.write(layer_index, slots, packed) - - def read_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - return self.mla_pool.read(layer_index, slots) - - def write_indexer_k(self, layer_index: int, slots: torch.Tensor, packed: torch.Tensor) -> None: - assert self.indexer_pool is not None - self.indexer_pool.write(layer_index, slots, packed) - - def read_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - assert self.indexer_pool is not None - return self.indexer_pool.read(layer_index, slots) +class DeepseekV4MemoryManager(MemoryManager): + """DeepSeek-V4 KV cache: 窗口 latent(全层) + c4/c128 压缩 latent(压实层) + c4 indexer-K。 -class DeepseekV4MemoryManager(Deepseek2MemoryManager): - """DeepSeek-V4 token-slot KV 管理(584B packed cache + bf16 workspace)。 + 与兄弟 manager 一致的 token-slot 设计;req 索引的表都在 DeepseekV4ReqManager。 - - dense/SWA latent: 主 ``kv_buffer`` 仍是 LightLLM 的 token-slot cache,不分页;物理格式改为 - SGLang/vLLM 的 ``fp8_ds_mla``: 448B NoPE fp8 + 64*2B RoPE bf16 + 7B scale + 1B pad = 584B。 - - c4_pool / c128_pool: 两个独立 ``_SubKvPool``(window 粒度,1-token 分配),compressed KV 同样 - 存 584B packed。c4 池附带 132B/token 的 packed indexer-K。 - - 读取时先用 torch reference dequant/gather 回 bf16 workspace,供现有 vLLM sparse FlashMLA wrapper - 消费;下一步可把这些 pack/dequant helper 替换成 fused/triton 版本。 - - 容量: 用闭式 ``get_cell_size()``(= 每个 dense token 在所有池上的 packed 总字节)让基类 - ``profile_size`` 直接得到 full_token = dense 池大小,再按 1/4、1/128 派生压缩池大小。 - - compressor 递归状态放 DeepseekV4ReqManager。 + - ``swa_pool``: 584B packed latent,所有层。池子小于 full token 空间;prep 阶段 + ``alloc_swa_prefill/decode`` 按**页**(128 槽,位置对齐: slot(p)=page_base+p%128)分配, + 映射记录到 ``full_to_swa_indexs``(以 full token 槽位为键)。出窗槽位由 DeepseekV4ReqManager + 在 prep 阶段批量惰性回收(``evict_swa``,页存活计数减到 0 才整页归还);full 槽位释放时 + ``free`` 级联回收对应 swa 槽,所以 radix 驱逐/请求释放/暂停无需任何额外协议。 + 页 allocator 触底时先走压力阀(radix 对 ref==0 节点回收)再 assert。 + 没有 ring buffer,prefill chunk 大小不受 sliding_window 限制。 + - ``c4_pool``/``c128_pool``: 压缩 latent,按 qwen3next 的层号压实手法只为压缩层建层; + c4 另带 packed indexer-K 池。槽位映射(``full_to_c4/c128_indexs``)以组末 token 的 full + 槽位为键(prep 阶段分配/scatter),``free`` 级联回收,与 swa 完全同构。 + - 写入走标准 operator 路径(``pack_mla_kv_to_cache``),内部为 triton packed writer; + torch codecs 保留为 ABI 的可执行规格(单测 oracle)。 """ operator_class = DeepseekV4MemOperator @@ -275,6 +155,7 @@ def __init__( indexer_head_dim: int = 128, max_request_num: Optional[int] = None, sliding_window: Optional[int] = None, + swa_extra_token_num: int = 0, always_copy=False, mem_fraction=0.9, ): @@ -290,15 +171,15 @@ def __init__( self.n_c4 = sum(1 for r in self.compress_rates if r == 4) self.n_c128 = sum(1 for r in self.compress_rates if r == 128) self.indexer_head_dim = indexer_head_dim - self.prefill_dtype = dtype - self.cache_dtype = torch.uint8 self.max_request_num = max_request_num self.sliding_window = sliding_window - self._pending_prefill_swa: Dict[int, Dict[str, torch.Tensor]] = {} + # 活跃窗口(max_request_num * sliding_window)之外的余量: 在途 prefill chunk 的瞬时占用 + # (出窗槽位要到下一次 prep 才回收) + radix cache 持有的窗口尾部。 + self.swa_extra_token_num = int(swa_extra_token_num) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) - self.layer_to_c4_idx: Dict[int, int] = {} - self.layer_to_c128_idx: Dict[int, int] = {} + self.layer_to_c4_idx = {} + self.layer_to_c128_idx = {} c4 = c128 = 0 for lid, r in enumerate(self.compress_rates): if r == 4: @@ -310,21 +191,56 @@ def __init__( super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) + # ------------------------------------------------------------------ sizing + def _swa_per_req_budget(self) -> int: + # 活跃请求保留 window + 一个 radix 页(req_manager._swa_retain_len: 让最近完成的 + # 128 边界的结尾页恒驻留,prompt cache 插入门才能放行),即 v5 §2 的「活跃窗口跨页 ≤2」。 + return int(self.sliding_window) + DSV4_SWA_PAGE_SIZE + def _planned_swa_size(self, full_size: int) -> int: + # swa 池按页分配(页 = 128 = sliding_window = radix 页),容量向上取整到整页。 if self.max_request_num is None or self.sliding_window is None: - return full_size - window_cap = max(1, int(self.max_request_num) * int(self.sliding_window)) - return max(1, min(full_size, window_cap)) + return _ceil_div(full_size, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE + cap = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num + cap = max(cap, int(full_size * DSV4_SWA_FULL_TOKENS_RATIO)) + cap = max(1, min(full_size, cap)) + return _ceil_div(cap, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE + + @staticmethod + def _slab_bytes_per_slot(page_size: int, data_bytes: int, scale_bytes: int, align_bytes: int = 1) -> float: + bytes_per_page = _ceil_div(page_size * (data_bytes + scale_bytes), align_bytes) * align_bytes + return bytes_per_page / page_size + + def _c4_state_bytes_per_swa_slot(self) -> float: + """c4 compressor state(attention + indexer,swa 页派生寻址)摊到每个 swa 槽的字节数。""" + if self.n_c4 == 0: + return 0.0 + per_page = DSV4_C4_STATE_RING * (4 * self.head_dim + 4 * self.indexer_head_dim) * 4 # fp32 + return per_page * self.n_c4 / DSV4_SWA_PAGE_SIZE + + def _swa_slot_bytes(self) -> float: + per_layer = self._slab_bytes_per_slot( + DSV4_SWA_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + return per_layer * self.layer_num + self._c4_state_bytes_per_swa_slot() - def _dense_cell_size(self): - return self.head_num * self.mla_bytes_per_token * self.layer_num + def _compressed_cell_size(self) -> float: + """每个 full token 摊到压缩池上的精确字节数(按 page-slab 对齐后)。""" + c4_latent = self._slab_bytes_per_slot( + DSV4_C4_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + c128_latent = self._slab_bytes_per_slot( + DSV4_C128_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + c4_indexer = self._slab_bytes_per_slot(DSV4_C4_PAGE_SIZE, self.indexer_head_dim, DSV4_INDEXER_SCALE_BYTES) + return (c4_latent + c4_indexer) * self.n_c4 / 4 + c128_latent * self.n_c128 / 128 - def _compressed_cell_size(self): - latent_bytes = self.head_num * self.mla_bytes_per_token - c4 = latent_bytes * self.n_c4 / 4 - c128 = latent_bytes * self.n_c128 / 128 - indexer = self.indexer_bytes_per_token * self.n_c4 / 4 - return c4 + c128 + indexer + def get_cell_size(self): + compressed = self._compressed_cell_size() + if self.size is None: + return self._swa_slot_bytes() + compressed + swa_ratio = self._planned_swa_size(self.size) / max(1, self.size) + return self._swa_slot_bytes() * swa_ratio + compressed def profile_size(self, mem_fraction): if self.size is not None: @@ -334,19 +250,26 @@ def profile_size(self, mem_fraction): world_size = dist.get_world_size() available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) available_bytes = available_memory * 1024 ** 3 - dense_cell = self._dense_cell_size() + swa_slot_bytes = self._swa_slot_bytes() compressed_cell = self._compressed_cell_size() if self.max_request_num is not None and self.sliding_window is not None and compressed_cell > 0: - swa_cap = max(1, int(self.max_request_num) * int(self.sliding_window)) - full_cell = dense_cell + compressed_cell - bytes_until_swa_cap = full_cell * swa_cap - if available_bytes <= bytes_until_swa_cap: + swa_budget = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num + full_cell = swa_slot_bytes + compressed_cell + if available_bytes <= full_cell * swa_budget: + # 小显存: full token 数还到不了 swa 预算,swa 池跟随 full token 数(每 token 一个 swa 槽)。 self.size = max(1, int(available_bytes / full_cell)) else: - self.size = max(1, int((available_bytes - dense_cell * swa_cap) / compressed_cell)) + size_budget = max(1, int((available_bytes - swa_slot_bytes * swa_budget) / compressed_cell)) + if size_budget * DSV4_SWA_FULL_TOKENS_RATIO > swa_budget: + # 比例下限生效(_planned_swa_size 会取 ratio*full),按该机制反解 full。 + self.size = max( + 1, int(available_bytes / (swa_slot_bytes * DSV4_SWA_FULL_TOKENS_RATIO + compressed_cell)) + ) + else: + self.size = size_budget else: - self.size = max(1, int(available_bytes / (dense_cell + compressed_cell))) + self.size = max(1, int(available_bytes / (swa_slot_bytes + compressed_cell))) if world_size > 1: tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") @@ -362,93 +285,350 @@ def profile_size(self, mem_fraction): logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" - f"{str((dense_cell + compressed_cell) / 1024 ** 2)} MB is the conservative size of one token kv cache\n" + f"{str(self.get_cell_size() / 1024 ** 2)} MB is the conservative size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) return - def get_cell_size(self): - dense = self._dense_cell_size() - compressed = self._compressed_cell_size() - if self.size is None: - return dense + compressed - swa_ratio = self._planned_swa_size(self.size) / max(1, self.size) - return dense * swa_ratio + compressed - + # ------------------------------------------------------------------ buffers def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): + rank_in_node = get_current_rank_in_node() + server = get_unique_server_name() + self.swa_size = self._planned_swa_size(size) - self.swa_pool = _PageSlabMlaPool( + assert self.swa_size % DSV4_SWA_PAGE_SIZE == 0 + self.swa_pool = PackedPagePool( size=self.swa_size, page_size=DSV4_SWA_PAGE_SIZE, layer_num=layer_num, - device="cuda", + data_bytes=DSV4_MLA_DATA_BYTES_PER_TOKEN, + scale_bytes=self.mla_scale_bytes, + align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES, ) - self.kv_buffer = self.swa_pool.kv_buffer - self._init_swa_mapping(size) - self._init_compressed_pools(size, head_num) - - def _init_swa_mapping(self, size): - rank_in_node = get_current_rank_in_node() - server = get_unique_server_name() - self.swa_allocator = KvCacheAllocator( - self.swa_size, - shared_name=f"{server}_dsv4_swa_can_use_token_num_{rank_in_node}", + # 注意: 该别名是 page 索引([layer, num_pages, bytes_per_page])而非 token 索引, + # 只允许 get_att_input_params 的消费者使用;token 索引语义的继承接口已显式 fence。 + self.kv_buffer = self.swa_pool.buffer + # 页粒度分配(页 = 128 槽,位置对齐): 槽位不变式 slot(p) = page_base + p%128。 + # swa_size 整页对齐 ⇒ HOLD 槽(swa_size)独占池子最后一个物理页,永不参与分配。 + self.swa_num_pages = self.swa_size // DSV4_SWA_PAGE_SIZE + self.swa_page_allocator = KvCacheAllocator( + self.swa_num_pages, shared_name=f"{server}_dsv4_swa_can_use_page_num_{rank_in_node}" ) + # 页存活计数 = 指向该页的有效 full_to_swa 行数;减到 0 归还 allocator(出窗逐 token + # 回收下,「部分出窗页」计数 > 0 自然受保护)。下标含 HOLD 页(只读不增减)。 + self.swa_page_live_count = torch.zeros((self.swa_pool.num_pages,), dtype=torch.int32, device="cuda") + # swa 压力阀(可选): 页 allocator 触底时回调(radix 对 ref==0 节点回收 swa 页), + # 由 backend 在 radix cache 创建后注入;assert 仍是最后防线。 + self._swa_pressure_valve = None self.full_to_swa_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") self.full_to_swa_indexs[size] = self.swa_pool.HOLD_TOKEN_MEMINDEX - if self.max_request_num is None or self.sliding_window is None: - self.req_to_swa_indexs = None - self.req_to_swa_full_indexs = None - return - - self.req_to_swa_indexs = torch.full( - (self.max_request_num + 1, self.sliding_window), - self.swa_pool.HOLD_TOKEN_MEMINDEX, - dtype=torch.int32, - device="cuda", - ) - self.req_to_swa_full_indexs = torch.full( - (self.max_request_num + 1, self.sliding_window), - -1, - dtype=torch.int32, - device="cuda", - ) - - def _init_compressed_pools(self, size, head_num): - rank_in_node = get_current_rank_in_node() - server = get_unique_server_name() - - self.c4_size = (size + 4 - 1) // 4 - self.c128_size = (size + 128 - 1) // 128 - self.c4_pool: Optional[_SubKvPool] = None - self.c128_pool: Optional[_SubKvPool] = None + self.c4_size = _ceil_div(size, 4) + self.c128_size = _ceil_div(size, 128) + self.c4_pool: Optional[PackedPagePool] = None + self.c4_indexer_pool: Optional[PackedPagePool] = None + self.c4_allocator: Optional[KvCacheAllocator] = None + self.c128_pool: Optional[PackedPagePool] = None + self.c128_allocator: Optional[KvCacheAllocator] = None + # 压缩槽映射: 键 = 组末 token(位置 (g+1)%ratio==0)的 full 槽位,值 = 压缩池槽位。 + # 与 full_to_swa_indexs 同构: radix 持有 full 槽 => 映射行存活,free 级联回收。 + self.full_to_c4_indexs: Optional[torch.Tensor] = None + self.full_to_c128_indexs: Optional[torch.Tensor] = None if self.n_c4 > 0: - self.c4_pool = _SubKvPool( + self.c4_pool = PackedPagePool( size=self.c4_size, page_size=DSV4_C4_PAGE_SIZE, layer_num=self.n_c4, - with_indexer=True, - shared_name=f"{server}_dsv4_c4_can_use_token_num_{rank_in_node}", + data_bytes=DSV4_MLA_DATA_BYTES_PER_TOKEN, + scale_bytes=self.mla_scale_bytes, + align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES, ) + self.c4_indexer_pool = PackedPagePool( + size=self.c4_size, + page_size=DSV4_C4_PAGE_SIZE, + layer_num=self.n_c4, + data_bytes=self.indexer_head_dim, + scale_bytes=DSV4_INDEXER_SCALE_BYTES, + ) + self.c4_allocator = KvCacheAllocator( + self.c4_size, shared_name=f"{server}_dsv4_c4_can_use_token_num_{rank_in_node}" + ) + self.full_to_c4_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") + self.full_to_c4_indexs[size] = self.c4_pool.HOLD_TOKEN_MEMINDEX + # c4 compressor 在途状态(attention + indexer): swa 页派生寻址(翻译③),随 swa 页 + # 生灭 -> radix 命中零拷贝续算。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), + # 取整到 ratio;末行哨兵 kv=0/score=-inf(KVAndScore.clear 语义),其余行由内核在 + # 组起点覆写,无需按页清零。last_dim = 2*coff*head_dim(overlap coff=2)。 + state_rows = self.swa_num_pages * DSV4_C4_STATE_RING + DSV4_C4_STATE_RING + 1 + state_rows = _ceil_div(state_rows, 4) * 4 + self.c4_state_buffer = torch.zeros( + (self.n_c4, state_rows, 4 * self.head_dim), dtype=torch.float32, device="cuda" + ) + self.c4_indexer_state_buffer = torch.zeros( + (self.n_c4, state_rows, 4 * self.indexer_head_dim), dtype=torch.float32, device="cuda" + ) + for buf in (self.c4_state_buffer, self.c4_indexer_state_buffer): + half = buf.shape[-1] // 2 + buf[:, -1, half:].fill_(float("-inf")) if self.n_c128 > 0: - self.c128_pool = _SubKvPool( + self.c128_pool = PackedPagePool( size=self.c128_size, page_size=DSV4_C128_PAGE_SIZE, layer_num=self.n_c128, - with_indexer=False, - shared_name=f"{server}_dsv4_c128_can_use_token_num_{rank_in_node}", + data_bytes=DSV4_MLA_DATA_BYTES_PER_TOKEN, + scale_bytes=self.mla_scale_bytes, + align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES, ) + self.c128_allocator = KvCacheAllocator( + self.c128_size, shared_name=f"{server}_dsv4_c128_can_use_token_num_{rank_in_node}" + ) + self.full_to_c128_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") + self.full_to_c128_indexs[size] = self.c128_pool.HOLD_TOKEN_MEMINDEX logger.info( - f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size} " + f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " f"packed_kv_bytes={self.mla_bytes_per_token} indexer_bytes={self.indexer_bytes_per_token}" ) + # ------------------------------------------------------------------ buffer accessors def get_att_input_params(self, layer_index: int): return self.swa_pool.get_layer_buffer(layer_index) + def _pool_and_local_layer(self, layer_index: int): + r = self.compress_rates[layer_index] + if r == 4: + return self.c4_pool, self.layer_to_c4_idx[layer_index] + if r == 128: + return self.c128_pool, self.layer_to_c128_idx[layer_index] + raise AssertionError(f"layer {layer_index} (rate {r}) 不是压缩层,没有压缩池") + + def get_compressed_kv_buffer(self, layer_index: int) -> torch.Tensor: + pool, local_layer = self._pool_and_local_layer(layer_index) + return pool.get_layer_buffer(local_layer) + + def get_indexer_k_buffer(self, layer_index: int) -> torch.Tensor: + assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" + return self.c4_indexer_pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]) + + def get_c4_state_buffer(self, layer_index: int) -> torch.Tensor: + assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 paged compressor state" + return self.c4_state_buffer[self.layer_to_c4_idx[layer_index]] + + def get_c4_indexer_state_buffer(self, layer_index: int) -> torch.Tensor: + assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 paged indexer state" + return self.c4_indexer_state_buffer[self.layer_to_c4_idx[layer_index]] + + # ------------------------------------------------------------------ swa slot lifecycle + def set_swa_pressure_valve(self, valve) -> None: + """valve(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点回收 swa)。""" + self._swa_pressure_valve = valve + return + + def _alloc_swa_pages(self, need_pages: int) -> torch.Tensor: + if need_pages > self.swa_page_allocator.can_use_mem_size and self._swa_pressure_valve is not None: + self._swa_pressure_valve(need_pages - self.swa_page_allocator.can_use_mem_size) + return self.swa_page_allocator.alloc(need_pages) + + def _count_swa_pages(self, swa_slots: torch.Tensor, delta: int) -> torch.Tensor: + """按 slot 所在页更新存活计数,返回触达的页(去重)。""" + pages = torch.div(swa_slots.long(), DSV4_SWA_PAGE_SIZE, rounding_mode="floor") + ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) + self.swa_page_live_count.index_add_(0, pages, ones) + return torch.unique(pages) + + def alloc_swa_prefill( + self, + b_req_idx: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_seq_len: torch.Tensor, + req_to_token_indexs: torch.Tensor, + ) -> None: + """prefill prep: 为各请求位置 [ready, seq) 的新 token 分配位置对齐的 swa 槽。 + + 槽位不变式: slot(p) = page_base(p 所在页) + p%128,page_base % 128 == 0。 + 续页(start 非整页,只可能是首页)的 base 从上一 token 的映射派生 + (full_to_swa[req_to_token[req, start-1]],该 token 必在保留窗内);其余页全新分配。 + radix 命中(ready 必 128 对齐)的借用方从全新页开始,与节点持有页天然不相交。 + 必须在 init_req_to_token_indexes 之后调用(scatter 目标经 req_to_token 行)。 + """ + page = DSV4_SWA_PAGE_SIZE + hold_req_id = self.max_request_num # padding 行的请求 id(req_manager.HOLD_REQUEST_ID) + req_list = b_req_idx.detach().cpu().tolist() + ready_list = b_ready_cache_len.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + + segs = [] # (req_idx, start, end, n_new_pages, has_cont_page) + total_new_pages = 0 + for req_idx, start, end in zip(req_list, ready_list, seq_list): + req_idx, start, end = int(req_idx), int(start), int(end) + if req_idx == hold_req_id or end <= start: + continue + first_new_page = _ceil_div(start, page) + n_new = max(0, (end - 1) // page - first_new_page + 1) + segs.append((req_idx, start, end, n_new, start % page != 0)) + total_new_pages += n_new + if not segs: + return + + new_pages = self._alloc_swa_pages(total_new_pages).cuda(non_blocking=True).long() if total_new_pages else None + page_cursor = 0 + for req_idx, start, end, n_new, has_cont in segs: + positions = torch.arange(start, end, dtype=torch.long, device="cuda") + page_local = torch.div(positions, page, rounding_mode="floor") - start // page + bases = torch.empty(((end - 1) // page - start // page + 1,), dtype=torch.long, device="cuda") + if has_cont: + prev_slot = int(self.full_to_swa_indexs[req_to_token_indexs[req_idx, start - 1].long()].item()) + # 续页不变式: 上一 token 必驻留(retain >= 2)且位置对齐(未来 resume/MTP 改动的哨兵)。 + assert prev_slot >= 0 and prev_slot % page == (start - 1) % page + bases[0] = prev_slot - (start - 1) % page + if n_new: + bases[1 if has_cont else 0 :] = new_pages[page_cursor : page_cursor + n_new] * page + page_cursor += n_new + slots = (bases[page_local] + positions % page).to(torch.int32) + self.full_to_swa_indexs[req_to_token_indexs[req_idx, start:end].long()] = slots + self._count_swa_pages(slots, 1) + return + + def alloc_swa_decode( + self, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + mem_indexes: torch.Tensor, + req_to_token_indexs: torch.Tensor, + ) -> None: + """decode prep: 本步 token(位置 seq-1)的 swa 槽。整页起点开新页,否则上一 token 槽 +1 + (位置对齐不变式保证同页连续)。scatter 目标用 mem_indexes(此刻 req_to_token 尚未写本步)。 + + 注意: 续槽从上一位置的映射派生,故同一请求的多行(MTP 多 token/步)在同一批内不支持 + (DSV4 启动参数已拒绝 MTP;支持需按步内顺序分段派生)。""" + page = DSV4_SWA_PAGE_SIZE + hold_req_id = self.max_request_num + req_list = b_req_idx.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + cont_rows, cont_prev_pos, new_rows = [], [], [] + for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): + req_idx, seq_len = int(req_idx), int(seq_len) + if req_idx == hold_req_id or seq_len <= 0: + continue + if (seq_len - 1) % page == 0: + new_rows.append(i) + else: + cont_rows.append(i) + cont_prev_pos.append(seq_len - 2) + mem_indexes = mem_indexes.cuda().long().reshape(-1) + if cont_rows: + req_rows = b_req_idx[cont_rows].long() + prev_full = req_to_token_indexs[req_rows, torch.tensor(cont_prev_pos, device="cuda")].long() + prev_slots = self.full_to_swa_indexs[prev_full] + # 续槽不变式哨兵: 上一位置必驻留(retain 覆盖)。prep 阶段本就有同步,代价可忽略。 + assert bool((prev_slots >= 0).all()) + slots = prev_slots + 1 + self.full_to_swa_indexs[mem_indexes[cont_rows]] = slots + self._count_swa_pages(slots, 1) + if new_rows: + pages = self._alloc_swa_pages(len(new_rows)).cuda(non_blocking=True).long() + slots = (pages * page).to(torch.int32) + self.full_to_swa_indexs[mem_indexes[new_rows]] = slots + self._count_swa_pages(slots, 1) + return + + def evict_swa(self, full_slots: torch.Tensor) -> None: + """回收 full 槽位对应的 swa 槽(出窗惰性回收 / free 级联 / 压力阀共用)。 + 未映射(-1)的槽位跳过;页计数减到 0 时整页归还 allocator。""" + if full_slots.numel() == 0: + return + full_slots = full_slots.cuda().long().reshape(-1) + full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) + if full_slots.numel() == 0: + return + swa_slots = self.full_to_swa_indexs[full_slots] + valid = swa_slots >= 0 + valid_slots = swa_slots[valid] + if valid_slots.numel() == 0: + return + self.full_to_swa_indexs[full_slots[valid]] = -1 + touched = self._count_swa_pages(valid_slots, -1) + empty = touched[self.swa_page_live_count[touched] == 0] + if empty.numel() > 0: + self.swa_page_allocator.free(empty.to(torch.int32)) + return + + def _evict_compress(self, full_slots: torch.Tensor, mapping: torch.Tensor, allocator: KvCacheAllocator) -> None: + full_slots = full_slots.cuda().long().reshape(-1) + # 去重: 同批重复槽会 gather 出重复的压缩槽 -> allocator 双重释放(free 已去重,直呼叫方防御)。 + full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) + if full_slots.numel() == 0: + return + slots = mapping[full_slots] + valid = slots >= 0 + valid_slots = slots[valid] + if valid_slots.numel() == 0: + return + allocator.free(valid_slots) + mapping[full_slots[valid]] = -1 + return + + def evict_c4(self, full_slots: torch.Tensor) -> None: + """回收 full 槽位(组末 token)映射的 c4 槽。非组末/未映射(-1)的槽位跳过。""" + if self.c4_allocator is None or full_slots.numel() == 0: + return + self._evict_compress(full_slots, self.full_to_c4_indexs, self.c4_allocator) + return + + def evict_c128(self, full_slots: torch.Tensor) -> None: + """回收 full 槽位(组末 token)映射的 c128 槽。非组末/未映射(-1)的槽位跳过。""" + if self.c128_allocator is None or full_slots.numel() == 0: + return + self._evict_compress(full_slots, self.full_to_c128_indexs, self.c128_allocator) + return + + # ------------------------------------------------------------------ alloc/free (cascade) + def free(self, free_index: Union[torch.Tensor, List[int]]) -> None: + """释放 full token 槽位,级联回收其 swa 槽与 c4/c128 压缩槽。radix 驱逐、请求释放/暂停都走这里。 + + 先对 full 槽去重: 同批重复槽位会让映射 gather 出重复的压缩/swa 槽,导致 allocator 双重释放。""" + if isinstance(free_index, list): + free_index = torch.tensor(free_index, dtype=torch.int64) + if free_index.numel() > 0: + free_index = torch.unique(free_index) + self.evict_swa(free_index) + self.evict_c4(free_index) + self.evict_c128(free_index) + super().free(free_index) + return + + def free_all(self): + super().free_all() + self.swa_page_allocator.free_all() + self.swa_page_live_count.zero_() + self.full_to_swa_indexs.fill_(-1) + self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX + if self.c4_allocator is not None: + self.c4_allocator.free_all() + self.full_to_c4_indexs.fill_(-1) + self.full_to_c4_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c4_pool.HOLD_TOKEN_MEMINDEX + if self.c128_allocator is not None: + self.c128_allocator.free_all() + self.full_to_c128_indexs.fill_(-1) + self.full_to_c128_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c128_pool.HOLD_TOKEN_MEMINDEX + return + + def alloc_c4(self, need_size) -> torch.Tensor: + return self.c4_allocator.alloc(need_size) + + def alloc_c128(self, need_size) -> torch.Tensor: + return self.c128_allocator.alloc(need_size) + + def free_c4(self, free_index) -> None: + self.c4_allocator.free(free_index) + + def free_c128(self, free_index) -> None: + self.c128_allocator.free(free_index) + + # ------------------------------------------------------------------ packed codecs (torch reference) + # 与 sglang/vllm 的 fp8_ds_mla 字节布局逐位对齐(ue8m0 幂次 scale)。这些 torch 实现是该 ABI 的 + # 可执行规格(单测 oracle,triton writer 与其逐字节对拍),不可删除。 def _pack_mla_kv(self, kv: torch.Tensor) -> torch.Tensor: kv = kv.reshape(-1, self.mla_head_dim) out = torch.empty((kv.shape[0], self.mla_bytes_per_token), dtype=torch.uint8, device=kv.device) @@ -500,7 +680,7 @@ def _pack_indexer_k(self, indexer_k: torch.Tensor) -> torch.Tensor: ) k_fp8 = torch.clamp(k_float / scale, -DSV4_FP8_E4M3_MAX, DSV4_FP8_E4M3_MAX).to(torch.float8_e4m3fn) out[:, : self.indexer_head_dim].copy_(k_fp8.view(dtype=torch.uint8)) - out[:, self.indexer_head_dim : self.indexer_bytes_per_token].copy_(scale.view(dtype=torch.uint8).reshape(-1, 4)) + out[:, self.indexer_head_dim :].copy_(scale.view(dtype=torch.uint8).reshape(-1, DSV4_INDEXER_SCALE_BYTES)) return out def _unpack_indexer_k(self, packed: torch.Tensor) -> torch.Tensor: @@ -508,483 +688,99 @@ def _unpack_indexer_k(self, packed: torch.Tensor) -> torch.Tensor: if packed.shape[0] == 0: return torch.empty((0, self.indexer_head_dim), dtype=self.dtype, device=packed.device) k_fp8 = packed[:, : self.indexer_head_dim].view(dtype=torch.float8_e4m3fn).float() - scale = packed[:, self.indexer_head_dim : self.indexer_bytes_per_token].view(dtype=torch.float32) + scale = packed[:, self.indexer_head_dim :].view(dtype=torch.float32) return (k_fp8 * scale).to(self.dtype) - def _identity_swa_slots(self, full_slots: torch.Tensor) -> torch.Tensor: - full_slots = full_slots.long() - valid = full_slots != self.HOLD_TOKEN_MEMINDEX - if valid.any() and int(full_slots[valid].max().item()) >= self.swa_size: - raise RuntimeError( - "DeepSeek-V4 SWA cache needs req_idx/positions for full token slots outside the SWA pool" - ) - swa_slots = torch.where( - valid, - full_slots, - torch.full_like(full_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), - ) - if valid.any(): - self.full_to_swa_indexs[full_slots[valid]] = swa_slots[valid].to(torch.int32) - return swa_slots - - def ensure_swa_slots(self, req_idx: int, positions: torch.Tensor, full_slots: torch.Tensor) -> torch.Tensor: - full_slots = full_slots.long().reshape(-1) - if full_slots.numel() == 0: - return full_slots - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: - return self._identity_swa_slots(full_slots) - - positions = positions.long().reshape(-1) - assert positions.numel() == full_slots.numel() - req_idx = int(req_idx) - out = torch.empty_like(full_slots, dtype=torch.long) - for i, (pos, full) in enumerate(zip(positions.tolist(), full_slots.tolist())): - if full == self.HOLD_TOKEN_MEMINDEX: - out[i] = self.swa_pool.HOLD_TOKEN_MEMINDEX - continue - - ring_pos = pos % self.sliding_window - old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) - old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) - if old_full == full and old_swa != self.swa_pool.HOLD_TOKEN_MEMINDEX: - swa = old_swa - elif old_swa != self.swa_pool.HOLD_TOKEN_MEMINDEX: - if old_full >= 0: - self.full_to_swa_indexs[old_full] = -1 - swa = old_swa - else: - swa = int(self.swa_allocator.alloc(1)[0].item()) - - self.req_to_swa_indexs[req_idx, ring_pos] = swa - self.req_to_swa_full_indexs[req_idx, ring_pos] = full - self.full_to_swa_indexs[full] = swa - out[i] = swa - return out - - def prepare_decode_swa_slots( - self, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, - mem_index: torch.Tensor, - ) -> None: - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: - return - - reqs = b_req_idx.detach().cpu().tolist() - seqs = b_seq_len.detach().cpu().tolist() - fulls = mem_index.detach().cpu().tolist() - hold = self.swa_pool.HOLD_TOKEN_MEMINDEX - for req_idx, seq_len, full in zip(reqs, seqs, fulls): - req_idx = int(req_idx) - full = int(full) - if req_idx == self.max_request_num or full == self.HOLD_TOKEN_MEMINDEX: - continue - ring_pos = (int(seq_len) - 1) % int(self.sliding_window) - old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) - old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) - if old_swa == hold: - old_swa = int(self.swa_allocator.alloc(1)[0].item()) - if old_full >= 0 and old_full != full: - self.full_to_swa_indexs[old_full] = -1 - self.req_to_swa_indexs[req_idx, ring_pos] = old_swa - self.req_to_swa_full_indexs[req_idx, ring_pos] = full - self.full_to_swa_indexs[full] = old_swa - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = hold - return - - def _reserve_prefill_swa_slots( - self, - req_idx: int, - positions: torch.Tensor, - full_slots: torch.Tensor, - ) -> Dict[str, torch.Tensor]: - full_slots = full_slots.long().reshape(-1) - positions = positions.long().reshape(-1) - assert positions.numel() == full_slots.numel() - - out = torch.empty_like(full_slots, dtype=torch.long) - ring_to_swa: Dict[int, int] = {} - ring_to_old_full: Dict[int, int] = {} - ring_to_final_full: Dict[int, int] = {} - hold = self.swa_pool.HOLD_TOKEN_MEMINDEX - - for i, (pos, full) in enumerate(zip(positions.tolist(), full_slots.tolist())): - if full == self.HOLD_TOKEN_MEMINDEX: - out[i] = hold - continue - - ring_pos = int(pos) % int(self.sliding_window) - swa = ring_to_swa.get(ring_pos) - if swa is None: - old_swa = int(self.req_to_swa_indexs[req_idx, ring_pos].item()) - old_full = int(self.req_to_swa_full_indexs[req_idx, ring_pos].item()) - if old_swa == hold: - old_swa = int(self.swa_allocator.alloc(1)[0].item()) - swa = old_swa - ring_to_swa[ring_pos] = swa - ring_to_old_full[ring_pos] = old_full - - ring_to_final_full[ring_pos] = int(full) - out[i] = swa - - rings = sorted(ring_to_final_full) - return { - "positions": positions.detach().clone(), - "full_slots": full_slots.detach().clone(), - "swa_slots": out.detach().clone(), - "commit_rings": torch.tensor(rings, dtype=torch.long, device=full_slots.device), - "commit_full_slots": torch.tensor( - [ring_to_final_full[r] for r in rings], - dtype=torch.long, - device=full_slots.device, - ), - "commit_swa_slots": torch.tensor( - [ring_to_swa[r] for r in rings], - dtype=torch.long, - device=full_slots.device, - ), - "commit_old_full_slots": torch.tensor( - [ring_to_old_full[r] for r in rings], - dtype=torch.long, - device=full_slots.device, - ), - } - - def prepare_prefill_swa_slots( - self, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, - b_ready_cache_len: torch.Tensor, - b_start_loc: torch.Tensor, - mem_index: torch.Tensor, - ) -> None: - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: - return - - self._pending_prefill_swa = {} - req_list = b_req_idx.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() - ready_list = b_ready_cache_len.detach().cpu().tolist() - start_list = b_start_loc.detach().cpu().tolist() - for req_idx, seq_len, ready_len, start_loc in zip(req_list, seq_list, ready_list, start_list): - token_num = int(seq_len) - int(ready_len) - if token_num <= 0: - continue - pos = torch.arange(int(ready_len), int(seq_len), dtype=torch.long, device=mem_index.device) - slots = mem_index[int(start_loc) : int(start_loc) + token_num] - self._pending_prefill_swa[int(req_idx)] = self._reserve_prefill_swa_slots(int(req_idx), pos, slots) - return - - def _get_pending_prefill_swa_slots( - self, - req_idx: int, - positions: torch.Tensor, - full_slots: torch.Tensor, - ) -> Optional[torch.Tensor]: - pending = self._pending_prefill_swa.get(int(req_idx)) - if pending is None: - return None - if pending["positions"].numel() != positions.numel(): - return None - if not torch.equal(pending["positions"].to(positions.device), positions.long().reshape(-1)): - return None - if not torch.equal(pending["full_slots"].to(full_slots.device), full_slots.long().reshape(-1)): - return None - return pending["swa_slots"].to(full_slots.device) - - def commit_prefill_swa_slots(self) -> None: - if not self._pending_prefill_swa: - return - for req_idx, pending in self._pending_prefill_swa.items(): - rings = pending["commit_rings"].to(self.req_to_swa_indexs.device) - if rings.numel() == 0: - continue - old_full = pending["commit_old_full_slots"].to(self.full_to_swa_indexs.device) - valid_old = old_full >= 0 - if valid_old.any(): - self.full_to_swa_indexs[old_full[valid_old].long()] = -1 - - full_slots = pending["commit_full_slots"].to(self.full_to_swa_indexs.device) - swa_slots = pending["commit_swa_slots"].to(self.full_to_swa_indexs.device) - self.req_to_swa_indexs[int(req_idx), rings] = swa_slots.to(torch.int32) - self.req_to_swa_full_indexs[int(req_idx), rings] = full_slots.to(torch.int32) - self.full_to_swa_indexs[full_slots.long()] = swa_slots.to(torch.int32) - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - self._pending_prefill_swa = {} - return - - def _swa_slots_from_full(self, full_slots: torch.Tensor) -> torch.Tensor: - full_slots = full_slots.long().reshape(-1) - if full_slots.numel() == 0: - return full_slots - mapped = self.full_to_swa_indexs[full_slots].long() - missing = mapped < 0 - if missing.any(): - if self.req_to_swa_indexs is not None: - bad = int(full_slots[missing][0].item()) - raise RuntimeError(f"DeepSeek-V4 dense KV for full token slot {bad} has been evicted from SWA cache") - fallback = full_slots[missing] - fallback_valid = fallback < self.swa_size - if fallback_valid.all(): - mapped[missing] = fallback - self.full_to_swa_indexs[fallback] = fallback.to(torch.int32) - else: - bad = int(fallback[~fallback_valid][0].item()) - raise RuntimeError(f"DeepSeek-V4 dense KV for full token slot {bad} has been evicted from SWA cache") - return mapped - - def free_swa_for_req(self, req_idx: int) -> None: - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: - return - req_idx = int(req_idx) - slots = self.req_to_swa_indexs[req_idx] - full_slots = self.req_to_swa_full_indexs[req_idx] - valid_swa = slots != self.swa_pool.HOLD_TOKEN_MEMINDEX - if valid_swa.any(): - free_slots = torch.unique(slots[valid_swa]).detach().cpu() - self.swa_allocator.free(free_slots) - valid_full = full_slots >= 0 - if valid_full.any(): - self.full_to_swa_indexs[full_slots[valid_full].long()] = -1 - self.req_to_swa_indexs[req_idx].fill_(self.swa_pool.HOLD_TOKEN_MEMINDEX) - self.req_to_swa_full_indexs[req_idx].fill_(-1) - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - - def snapshot_swa_for_prompt_cache(self, req_idx: int, cache_len: int, full_slots: torch.Tensor): - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None or cache_len <= 0: - return None - tail_start = max(0, int(cache_len) - int(self.sliding_window)) - full_slots = full_slots[tail_start:cache_len].long().to(self.kv_buffer.device) - if full_slots.numel() == 0: - return None - swa_slots = self.full_to_swa_indexs[full_slots].long() - if (swa_slots < 0).any(): - bad = int(full_slots[swa_slots < 0][0].item()) - raise RuntimeError(f"DeepSeek-V4 prompt cache cannot snapshot evicted SWA full slot {bad}") - return { - "positions": torch.arange(tail_start, cache_len, dtype=torch.int64, device="cpu"), - "full_slots": full_slots.detach().cpu(), - "swa_slots": swa_slots.detach().cpu(), - } - - def clone_swa_for_prompt_cache(self, req_idx: int, cache_len: int, full_slots: torch.Tensor): - payload = self.snapshot_swa_for_prompt_cache(req_idx, cache_len, full_slots) - if payload is None: - return None - - src_slots = payload["swa_slots"].long().to(self.kv_buffer.device) - dst_slots = self.swa_allocator.alloc(src_slots.numel()).long().to(self.kv_buffer.device) - for layer_idx in range(self.layer_num): - self.swa_pool.write(layer_idx, dst_slots, self.swa_pool.read(layer_idx, src_slots)) - payload["swa_slots"] = dst_slots.detach().cpu() - return payload - - def detach_swa_for_prompt_cache(self, req_idx: int, swa_payload) -> None: - if ( - swa_payload is None - or self.req_to_swa_indexs is None - or self.req_to_swa_full_indexs is None - or len(swa_payload["positions"]) == 0 - ): - return - req_idx = int(req_idx) - positions = swa_payload["positions"].tolist() - full_slots = swa_payload["full_slots"].tolist() - swa_slots = swa_payload["swa_slots"].tolist() - for pos, full, swa in zip(positions, full_slots, swa_slots): - ring_pos = int(pos) % int(self.sliding_window) - if int(self.req_to_swa_indexs[req_idx, ring_pos].item()) == int(swa) and int( - self.req_to_swa_full_indexs[req_idx, ring_pos].item() - ) == int(full): - self.req_to_swa_indexs[req_idx, ring_pos] = self.swa_pool.HOLD_TOKEN_MEMINDEX - self.req_to_swa_full_indexs[req_idx, ring_pos] = -1 - return - - def restore_swa_from_prompt_cache(self, swa_payload) -> None: - if swa_payload is None or len(swa_payload["full_slots"]) == 0: - return - full_slots = swa_payload["full_slots"].long().to(self.kv_buffer.device) - swa_slots = swa_payload["swa_slots"].long().to(self.kv_buffer.device) - self.full_to_swa_indexs[full_slots] = swa_slots.to(torch.int32) - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - return - - def free_swa_prompt_cache(self, swa_payload) -> None: - if swa_payload is None or len(swa_payload["swa_slots"]) == 0: - return - swa_slots = torch.unique(swa_payload["swa_slots"].long()).detach().cpu() - self.swa_allocator.free(swa_slots) - full_slots = swa_payload["full_slots"].long().to(self.kv_buffer.device) - mapped = self.full_to_swa_indexs[full_slots].long() - expected = swa_payload["swa_slots"].long().to(self.kv_buffer.device) - same = mapped == expected - if same.any(): - self.full_to_swa_indexs[full_slots[same]] = -1 - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - return - - def _keep_last_swa_writes(self, swa_slots: torch.Tensor, packed: torch.Tensor): - """Drop duplicate SWA writes generated by long prefill ring reuse.""" - if swa_slots.numel() <= 1: - return swa_slots, packed - - slots_cpu = swa_slots.detach().cpu().tolist() - seen = set() - keep = [] - hold = self.swa_pool.HOLD_TOKEN_MEMINDEX - for i in range(len(slots_cpu) - 1, -1, -1): - slot = int(slots_cpu[i]) - if slot == hold or slot in seen: - continue - seen.add(slot) - keep.append(i) - keep.reverse() - if len(keep) == len(slots_cpu): - return swa_slots, packed - if not keep: - return swa_slots[:0], packed[:0] - keep_index = torch.tensor(keep, dtype=torch.long, device=swa_slots.device) - return swa_slots.index_select(0, keep_index), packed.index_select(0, keep_index) - - def pack_mla_kv_to_cache( - self, - layer_index: int, - mem_index: torch.Tensor, - kv: torch.Tensor, - req_idx: Optional[int] = None, - positions: Optional[torch.Tensor] = None, - ): + # ------------------------------------------------------------------ cache write paths + def pack_mla_kv_to_cache(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): + """标准 operator 写入路径。要求本步已对 mem_index 调过 ``alloc_swa``(prep 阶段); + HOLD/padding 槽位映射到 swa HOLD 槽,写入无害。""" if kv.shape[0] == 0: return - packed = self._pack_mla_kv(kv) - if req_idx is None or positions is None: - swa_slots = self._identity_swa_slots(mem_index).to(kv.device) - else: - pending_slots = self._get_pending_prefill_swa_slots(req_idx, positions, mem_index) - if pending_slots is None: - swa_slots = self.ensure_swa_slots(req_idx, positions, mem_index).to(kv.device) - else: - swa_slots = pending_slots.to(kv.device) - swa_slots, packed = self._keep_last_swa_writes(swa_slots, packed) - if swa_slots.numel() == 0: - return - self.swa_pool.write(layer_index, swa_slots, packed) - - def pack_decode_mla_kv_to_cache( - self, - layer_index: int, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, - mem_index: torch.Tensor, - kv: torch.Tensor, - ): - if kv.shape[0] == 0: - return - packed = self._pack_mla_kv(kv) - if self.req_to_swa_indexs is None or self.req_to_swa_full_indexs is None: - swa_slots = self._identity_swa_slots(mem_index).to(kv.device) - else: - req = b_req_idx.long() - ring = ((b_seq_len.long() - 1) % int(self.sliding_window)).long() - swa_slots = self.req_to_swa_indexs[req, ring].long() - - old_full = self.req_to_swa_full_indexs[req, ring].long() - full_slots = mem_index.long() - old_full = torch.where(old_full >= 0, old_full, full_slots) - self.full_to_swa_indexs[old_full] = torch.full( - old_full.shape, - -1, - dtype=self.full_to_swa_indexs.dtype, - device=old_full.device, - ) - - self.req_to_swa_full_indexs[req, ring] = full_slots.to(torch.int32) - self.full_to_swa_indexs[full_slots] = swa_slots.to(torch.int32) - self.swa_pool.write(layer_index, swa_slots.to(kv.device), packed) + from lightllm.models.deepseek_v4.triton_kernel.destindex_copy_kv_flashmla_dsv4 import ( + destindex_copy_kv_flashmla_dsv4, + ) - def gather_mla_kv_from_swa_slots(self, layer_index: int, swa_slots: torch.Tensor) -> torch.Tensor: - return self._unpack_mla_kv(self.swa_pool.read(layer_index, swa_slots.to(self.kv_buffer.device))) + swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] + destindex_copy_kv_flashmla_dsv4( + kv.reshape(-1, self.mla_head_dim), + swa_slots, + self.swa_pool.get_layer_buffer(layer_index), + self.swa_pool.page_size, + ) + return def pack_compressed_kv_to_cache(self, layer_index: int, slots: torch.Tensor, comp: torch.Tensor): if comp.shape[0] == 0: return + from lightllm.models.deepseek_v4.triton_kernel.destindex_copy_kv_flashmla_dsv4 import ( + destindex_copy_kv_flashmla_dsv4, + ) + pool, local_layer = self._pool_and_local_layer(layer_index) - pool.write_kv(local_layer, slots.to(comp.device), self._pack_mla_kv(comp)) + destindex_copy_kv_flashmla_dsv4( + comp.reshape(-1, self.mla_head_dim), + slots.to(comp.device), + pool.get_layer_buffer(local_layer), + pool.page_size, + ) - def pack_c4_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer_k: torch.Tensor): + def pack_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer_k: torch.Tensor): if indexer_k.shape[0] == 0: return - pool, local_layer = self._pool_and_local_layer(layer_index) - pool.write_indexer_k(local_layer, slots.to(indexer_k.device), self._pack_indexer_k(indexer_k)) - - def gather_mla_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - if slots.numel() == 0: - return torch.empty((0, self.mla_head_dim), dtype=self.dtype, device=self.kv_buffer.device) - swa_slots = self._swa_slots_from_full(slots).to(self.kv_buffer.device) - return self._unpack_mla_kv(self.swa_pool.read(layer_index, swa_slots)) - - def gather_compressed_kv(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - if slots.numel() == 0: - return torch.empty((0, self.mla_head_dim), dtype=self.dtype, device=self.kv_buffer.device) - pool, local_layer = self._pool_and_local_layer(layer_index) - return self._unpack_mla_kv(pool.read_kv(local_layer, slots.to(self.kv_buffer.device))) - - def gather_c4_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - if slots.numel() == 0: - return torch.empty( - (0, self.indexer_head_dim), - dtype=self.dtype, - device=self.kv_buffer.device, - ) - pool, local_layer = self._pool_and_local_layer(layer_index) - return self._unpack_indexer_k(pool.read_indexer_k(local_layer, slots.to(self.kv_buffer.device))) - - def _pool_and_local_layer(self, layer_index: int): - r = self.compress_rates[layer_index] - if r == 4: - return self.c4_pool, self.layer_to_c4_idx[layer_index] - if r == 128: - return self.c128_pool, self.layer_to_c128_idx[layer_index] - raise AssertionError(f"layer {layer_index} (rate {r}) 不是压缩层,没有压缩池") + assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" + from lightllm.models.deepseek_v4.triton_kernel.destindex_copy_indexer_k_dsv4 import ( + destindex_copy_indexer_k_dsv4, + ) - def get_compressed_kv_buffer(self, layer_index: int) -> torch.Tensor: - pool, local_layer = self._pool_and_local_layer(layer_index) - return pool.get_kv_buffer(local_layer) + destindex_copy_indexer_k_dsv4( + indexer_k.reshape(-1, self.indexer_head_dim), + slots.to(indexer_k.device), + self.c4_indexer_pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]), + self.c4_indexer_pool.page_size, + ) - def get_compressed_indexer_k_buffer(self, layer_index: int) -> torch.Tensor: + def gather_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: + """反量化 gather c4 indexer-K: slots [N](c4 槽位,HOLD 合法) -> [N, indexer_head_dim] bf16。 + indexer top-k 打分用(纯张量操作,cuda-graph 安全)。""" assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" - return self.c4_pool.get_index_k_buffer(self.layer_to_c4_idx[layer_index]) + pool = self.c4_indexer_pool + flat = pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]).view(-1) + data_offsets, scale_offsets = pool._loc_offsets(slots.reshape(-1)) + data_range = torch.arange(pool.data_bytes_per_token, device=flat.device) + scale_range = torch.arange(pool.scale_bytes_per_token, device=flat.device) + k_fp8 = flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)].view(torch.float8_e4m3fn) + scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)].contiguous().view(torch.float32) + return (k_fp8.float() * scale).to(torch.bfloat16) + + # ------------------------------------------------------------------ fenced inherited APIs + # kv_buffer 是 page 索引的 uint8 slab,基类按 token 索引读写的接口会静默写坏数据,显式拦截。 + def get_index_kv_buffer(self, index): + raise NotImplementedError("DeepSeek-V4 packed page-slab cache does not support token-indexed kv_buffer io") + + def load_index_kv_buffer(self, index, load_tensor_dict): + raise NotImplementedError("DeepSeek-V4 packed page-slab cache does not support token-indexed kv_buffer io") - def alloc_c4(self, need_size) -> torch.Tensor: - return self.c4_pool.alloc(need_size) + def alloc_kv_move_buffer(self, max_req_total_len): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def alloc_c128(self, need_size) -> torch.Tensor: - return self.c128_pool.alloc(need_size) + def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def free_c4(self, free_index) -> None: - self.c4_pool.free(free_index) + def write_mem_to_page_kv_move_buffer(self, *args, **kwargs): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def free_c128(self, free_index) -> None: - self.c128_pool.free(free_index) + def read_page_kv_move_buffer_to_mem(self, *args, **kwargs): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def free_all(self): - super().free_all() - if hasattr(self, "swa_allocator"): - self.swa_allocator.free_all() - if hasattr(self, "full_to_swa_indexs"): - self.full_to_swa_indexs.fill_(-1) - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - if getattr(self, "req_to_swa_indexs", None) is not None: - self.req_to_swa_indexs.fill_(self.swa_pool.HOLD_TOKEN_MEMINDEX) - self.req_to_swa_full_indexs.fill_(-1) - self._pending_prefill_swa = {} - if self.c4_pool is not None: - self.c4_pool.free_all() - if self.c128_pool is not None: - self.c128_pool.free_all() + def send_to_decode_node(self, *args, **kwargs): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def alloc_kv_move_buffer(self, max_req_total_len): + def receive_from_prefill_node(self, *args, **kwargs): raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - raise NotImplementedError("DeepSeek-V4 packed/composite paged KV transfer is not implemented") + def send_to_decode_node_p2p(self, *args, **kwargs): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") + + def receive_from_prefill_node_p2p(self, *args, **kwargs): + raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") diff --git a/lightllm/common/quantization/__init__.py b/lightllm/common/quantization/__init__.py index 1c5a9c09d3..4db55e1555 100644 --- a/lightllm/common/quantization/__init__.py +++ b/lightllm/common/quantization/__init__.py @@ -14,6 +14,7 @@ EXPERT_DTYPE_TO_QUANT_TYPE = { "fp8": "deepgemm-fp8w8a8-b128", "fp4": "deepgemm-fp4fp8-b32", + "mxfp4": "marlin-mxfp4w4a16-b32", } SUPPORTED_EXPERT_DTYPES = tuple(EXPERT_DTYPE_TO_QUANT_TYPE) @@ -64,13 +65,13 @@ def _mapping_quant_method(self): logger.info(f"select fp8w8a8-b128 quant way: {self.quant_type}") # fp8 量化下,部分 MoE 模型(如 DeepSeek-V4),可以单独声明 expert 权重精度, - # 按其值给 fused_moe 选用对应的 deepgemm 量化方法。 + # 按其值给 fused_moe 选用对应的量化方法。 expert_dtype = self.expert_dtype or self.network_config_.get("expert_dtype", None) if expert_dtype is None: return - if expert_dtype == "fp4" and self.network_config_.get("model_type") == "deepseek_v4" and not is_sm100_gpu(): - logger.info("skip generic fused_moe quant mapping for DeepSeek-V4 fp4 experts on non-SM100 GPUs") - return + # DeepSeek-V4 的 fp4 发布版自带预打包 MXFP4 专家。 + if expert_dtype == "fp4" and self.network_config_.get("model_type") == "deepseek_v4": + expert_dtype = "mxfp4" target = self._get_expert_quant_type(expert_dtype) for layer_num in range(self.layer_num): if self.expert_dtype is not None: diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index ec1ee90fd4..677d3b7dd7 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -198,6 +198,84 @@ def _create_weight( return mm_param, mm_param_list +@QUANTMETHODS.register(["marlin-mxfp4w4a16-b32"], platform="cuda") +class MXFP4MoEQuantizationMethod(QuantizationMethod): + def __init__(self): + super().__init__() + self.block_size = 32 + self.weight_suffix = "weight" + self.weight_zero_point_suffix = None + self.weight_scale_suffix = "scale" + self.has_weight_scale = True + self.has_weight_zero_point = False + + @property + def method_name(self): + return "marlin-mxfp4w4a16-b32" + + def quantize(self, weight: torch.Tensor, output: WeightPack): + raise NotImplementedError("marlin-mxfp4w4a16-b32 only loads pre-packed MXFP4 expert weights") + + def apply( + self, + input_tensor: torch.Tensor, + weight_pack: "WeightPack", + out: Optional[torch.Tensor] = None, + workspace: Optional[torch.Tensor] = None, + use_custom_tensor_mananger: bool = True, + bias: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + raise NotImplementedError("marlin-mxfp4w4a16-b32 is only implemented for fused MoE expert weights") + + def _create_weight( + self, out_dims: Union[int, List[int]], in_dim: int, dtype: torch.dtype, device_id: int, num_experts: int = 1 + ) -> Tuple[WeightPack, List[WeightPack]]: + out_dim = sum(out_dims) if isinstance(out_dims, list) else out_dims + assert in_dim % 2 == 0, "MXFP4 packed weight requires even input dimension" + assert in_dim % self.block_size == 0, "MXFP4 scale dimension must be divisible by block_size" + expert_prefix = (num_experts,) if num_experts > 1 else () + weight = torch.empty(expert_prefix + (out_dim, in_dim // 2), dtype=torch.int8, device="cpu") + weight_scale = torch.empty( + expert_prefix + (out_dim, in_dim // self.block_size), dtype=torch.float8_e8m0fnu, device="cpu" + ) + mm_param = WeightPack(weight=weight, weight_scale=weight_scale) + mm_param_list = self._split_weight_pack( + mm_param, + weight_out_dims=out_dims, + weight_split_dim=-2, + weight_scale_out_dims=out_dims, + weight_scale_split_dim=-2, + ) + return mm_param, mm_param_list + + def finalize_moe_weight(self, moe_weight): + try: + from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( + prepare_moe_mxfp4_layer_for_marlin, + ) + except Exception as e: + raise RuntimeError(f"marlin-mxfp4w4a16-b32 requires vLLM MXFP4 packing utilities, error={repr(e)}") from e + + class _MXFP4Layer: + pass + + device = torch.device("cuda", moe_weight.device_id_) + layer = _MXFP4Layer() + layer.params_dtype = moe_weight.data_type_ + w13 = moe_weight.w13.weight.view(torch.uint8).to(device=device, non_blocking=True).contiguous() + w2 = moe_weight.w2.weight.view(torch.uint8).to(device=device, non_blocking=True).contiguous() + w13_scale = moe_weight.w13.weight_scale.to(device=device, non_blocking=True).contiguous() + w2_scale = moe_weight.w2.weight_scale.to(device=device, non_blocking=True).contiguous() + ( + moe_weight.w13.weight, + moe_weight.w2.weight, + moe_weight.w13.weight_scale, + moe_weight.w2.weight_scale, + _, + _, + ) = prepare_moe_mxfp4_layer_for_marlin(layer, w13, w2, w13_scale, w2_scale, None, None) + + def _deepgemm_fp8_nt(a_tuple, b_tuple, out): if HAS_DEEPGEMM: if hasattr(deep_gemm, "gemm_fp8_fp8_bf16_nt"): diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 7b56129c3f..24b39b71ce 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -26,14 +26,18 @@ @dataclass class DeepseekV4PromptCachePayload: + """prompt cache 载荷: 只剩 swa 按页有效性 bitmap。 + + 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 + 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4 compressor 状态以 swa 页派生 + 寻址(随 swa 页生灭,命中零拷贝续算);c128 状态在 128 边界自然归零,无需恢复。 + + * ``swa_page_valid``: cpu bool [cache_len // page],插入时按当下 full_to_swa 映射写定 + (页内 128 个映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 128 边界, + swa 压力阀回收节点页时清零。""" + cache_len: int - c4_slots: Optional[torch.Tensor] = None - c128_slots: Optional[torch.Tensor] = None - c4_state: Optional[torch.Tensor] = None - c4_state_pool: Optional[torch.Tensor] = None - c4_indexer_state: Optional[torch.Tensor] = None - c4_indexer_state_pool: Optional[torch.Tensor] = None - swa: Optional[dict] = None + swa_page_valid: Optional[torch.Tensor] = None class DeepseekV4PromptCacheValueOps: @@ -47,9 +51,29 @@ def concat(self, payloads: List[DeepseekV4PromptCachePayload]): return self.req_manager.concat_prompt_cache_payloads(payloads) def free(self, payload: DeepseekV4PromptCachePayload): - self.req_manager.free_prompt_cache_payload(payload) + # 槽位资源全部由 mem_manager.free(full_slots) 级联回收,载荷本身没有需要释放的资源。 return + def invalidate_swa_pages(self, payload: DeepseekV4PromptCachePayload) -> None: + """swa 压力阀回收了该节点的 swa 页后清 bitmap: 后续命中按缩短语义裁剪,不会复活。""" + if payload is not None and payload.swa_page_valid is not None: + payload.swa_page_valid.fill_(False) + return + + def valid_match_length(self, payload: Optional[DeepseekV4PromptCachePayload], natural_len: int) -> int: + """radix 匹配裁剪: 返回 <= natural_len 的最大 128 边界 L',使结尾页(bitmap[L'/128-1])有效。 + + 有效性可能非单调(owner 生前从左驱逐、后续阀从尾回收),按候选边界回查 bitmap; + 中段 invalid 页不挡更靠后的有效命中(注意力只回看最后一个窗口)。""" + page = self.req_manager.get_prompt_cache_page_size() + if payload is None or payload.swa_page_valid is None: + return 0 + n_pages = min(natural_len // page, int(payload.swa_page_valid.numel())) + valid_idx = torch.nonzero(payload.swa_page_valid[:n_pages]) + if valid_idx.numel() == 0: + return 0 + return (int(valid_idx[-1]) + 1) * page + class _ReqNode: def __init__(self, index): @@ -334,21 +358,23 @@ def copy_small_page_buffer_to_linear_att_state( class DeepseekV4ReqManager(ReqManager): - """DeepSeek-V4 的请求级管理(锁定决策: SWA 全历史 + 不分页)。 - - 在基类 ReqManager 之上补三类 V4 专有的 per-request 结构。该对象在 mem manager profile 前创建, - 所以初始化只依赖 config 派生出的 compress_rates/head_dim/indexer_head_dim;真实 mem_manager - 会在 `_init_mem_manager()` 后通过 `bind_mem_manager()` 接入。 - - * ``req_to_c4_indexs`` / ``req_to_c128_indexs`` —— (req, 窗口下标) -> 压缩池槽位。 - 窗口下标 = position // compress_rate;窗口关闭时由 layer-infer 写入,attention 读取前 - n_windows 列即该 req 的全部压缩条目槽。未填充列为 0(不会被读到,语义同 req_to_token_indexs)。 - * ``req_to_c4_state`` / ``req_to_c128_state`` / ``req_to_c4_indexer_state`` —— compressor 的 - “在途窗口”累加状态(per req、per 压缩层),fp32。形状为 - ``(kv_or_score, coff * ratio, coff * dim)``; c4 因 Ca/Cb overlap 取 ``coff=2``, - c128 取 ``coff=1``。score 初始化为 ``-inf``,与官方 reference compressor 的 - ``kv_state``/``score_state`` 对齐。 - * entry_count 不另存:= position // compress_rate,可由序列长度推出。 + """DeepSeek-V4 的请求级管理。 + + 在基类 ReqManager 之上补 V4 专有的 per-request 结构。该对象在 mem manager profile 前创建, + 所以初始化只依赖 config 派生出的 compress_rates/head_dim/indexer_head_dim/sliding_window; + 真实 mem_manager 会在 `_init_mem_manager()` 后通过 `bind_mem_manager()` 接入。 + + * 压缩槽位不在本类: ``full_to_c4/c128_indexs``(mem manager)以组末 token 的 full 槽位为键。 + 本类只负责 prep 阶段的分配与 scatter(``prepare_prefill_compress_slots`` / + ``prepare_decode_compress_slots``)——必须先于 attention metadata 构建/图捕获; + 条目内容由 layer-infer 的 compressor 前向写入。 + * ``req_to_c128_state_pool`` —— c128 compressor 的在途状态(per req、per c128 层)。 + c128 在线聚合在 128 边界自然归零(命中边界必 128 对齐),无缓存常驻需求,保持 req 键控。 + c4 状态(跨边界 overlap)在 mem manager 的 swa 页派生池,随页生灭,命中零拷贝续算。 + * SWA 槽位分配/出窗回收(``prepare_prefill_swa`` / ``prepare_decode_swa``): 每步 prep 阶段 + 为新 token 调 mem_manager.alloc_swa,并按 per-req 水位线(``_swa_evict_marks``)惰性回收 + 已出窗位置的 swa 槽。水位线首次置为该请求首个 chunk 的 ready_cache_len(radix 共享前缀 + 的边界),因此共享前缀的 swa 槽永远不会被本请求回收(归 radix 经 mem_manager.free 级联释放)。 """ def __init__( @@ -359,10 +385,24 @@ def __init__( compress_rates: Optional[List[int]] = None, head_dim: Optional[int] = None, indexer_head_dim: Optional[int] = None, + sliding_window: Optional[int] = None, ): super().__init__(max_request_num, max_sequence_length, mem_manager) self.mem_manager = mem_manager + if mem_manager is not None: + if compress_rates is None: + compress_rates = mem_manager.compress_rates + if head_dim is None: + head_dim = mem_manager.head_dim + if indexer_head_dim is None: + indexer_head_dim = mem_manager.indexer_head_dim + if sliding_window is None: + sliding_window = mem_manager.sliding_window + self.sliding_window = sliding_window + # 出窗回收水位线: -1 表示该 req 尚未见过 prefill chunk(首个 chunk 的 ready_cache_len + # 即共享前缀边界,作为永不下探的回收下界)。 + self._swa_evict_marks = [-1 for _ in range(max_request_num + 1)] self.compress_rates = list(compress_rates) self.n_c4 = sum(1 for r in self.compress_rates if r == 4) self.n_c128 = sum(1 for r in self.compress_rates if r == 128) @@ -379,60 +419,16 @@ def __init__( self.layer_to_c128_idx[lid] = c128 c128 += 1 - # (req, 窗口) -> 压缩槽。列数取 ceil(max_seq / ratio) 留足余量。 - c4_windows = (max_sequence_length + 4 - 1) // 4 - c128_windows = (max_sequence_length + 128 - 1) // 128 - self.req_to_c4_indexs = torch.zeros((max_request_num + 1, c4_windows), dtype=torch.int32, device="cuda") - self.req_to_c128_indexs = torch.zeros((max_request_num + 1, c128_windows), dtype=torch.int32, device="cuda") - self._c4_entry_counts = [0 for _ in range(max_request_num + 1)] - self._c128_entry_counts = [0 for _ in range(max_request_num + 1)] - - # compressor 在途窗口累加状态(fp32): [kv_or_score, coff * ratio, coff * dim]. - state_dtype = torch.float32 - self.req_to_c4_state = LayerCache( - size=max_request_num + 1, - dtype=state_dtype, - shape=(2, 8, 2 * head_dim), - layer_num=self.n_c4, - device="cuda", - ) - self.req_to_c128_state = LayerCache( - size=max_request_num + 1, - dtype=state_dtype, - shape=(2, 128, head_dim), - layer_num=self.n_c128, - device="cuda", - ) - self.req_to_c4_indexer_state = LayerCache( - size=max_request_num + 1, - dtype=state_dtype, - shape=(2, 8, 2 * indexer_head_dim), - layer_num=self.n_c4, - device="cuda", - ) - self.req_to_c4_state_pool = LayerCache( - size=max_request_num + 1, - dtype=state_dtype, - shape=(1, 8, 4 * head_dim), - layer_num=self.n_c4, - device="cuda", - ) + # c128 compressor 在途状态(fp32): 在线聚合在 128 边界自然归零(命中边界必 128 对齐), + # 无缓存常驻需求,保持 req 键控。c4 状态(有跨边界 overlap)在 mem manager 的 + # swa 页派生池(c4_state_buffer / c4_indexer_state_buffer)。 self.req_to_c128_state_pool = LayerCache( size=max_request_num + 1, - dtype=state_dtype, + dtype=torch.float32, shape=(1, 128, 2 * head_dim), layer_num=self.n_c128, device="cuda", ) - self.req_to_c4_indexer_state_pool = LayerCache( - size=max_request_num + 1, - dtype=state_dtype, - shape=(1, 8, 4 * indexer_head_dim), - layer_num=self.n_c4, - device="cuda", - ) - self._runtime_states = [{} for _ in range(max_request_num + 1)] - self._init_all_score_state() return def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): @@ -440,22 +436,92 @@ def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): assert self.compress_rates == mem_manager.compress_rates assert self.head_dim == mem_manager.head_dim assert self.indexer_head_dim == mem_manager.indexer_head_dim + if self.sliding_window is None: + self.sliding_window = mem_manager.sliding_window + else: + assert mem_manager.sliding_window is None or self.sliding_window == mem_manager.sliding_window self.mem_manager = mem_manager return - def _init_all_score_state(self): - if self.n_c4 > 0: - self.req_to_c4_state.buffer[:, :, 1, ...].fill_(float("-inf")) - self.req_to_c4_indexer_state.buffer[:, :, 1, ...].fill_(float("-inf")) - if self.n_c128 > 0: - self.req_to_c128_state.buffer[:, :, 1, ...].fill_(float("-inf")) + # ------------------------------------------------------------------ swa slot prep (per step) + def _swa_retain_len(self) -> int: + """出窗回收的保留长度 = window + 一个 radix 页。 + + 多留一页使「最近一个完成的 128 边界」的结尾页恒驻留: prompt cache 只能在 floor(cur/128) + 边界入树(radix page=128),若回收只留 window,则任何非对齐时刻该边界的结尾页都已被 + 部分回收,插入门会把所有插入裁到 0(prompt cache 形同虚设)。预算即 v5 §2 的每请求 + 「活跃窗口跨页 ≤2」。驻留证明要求 window >= page-1(DSV4 实际 window == page == 128)。""" + return int(self.sliding_window) + self.get_prompt_cache_page_size() + + def prepare_prefill_swa( + self, + b_req_idx: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_seq_len: torch.Tensor, + ) -> None: + """prefill prep: 为本 chunk 全部新 token(位置 [ready, seq))分配位置对齐的 swa 槽, + 并回收已出窗位置的槽。 + + 本 chunk 起点 L = ready_cache_len,首个新 token(位置 L)的窗口是 [L-W+1, L];回收 + 边界再额外保留一个 radix 页(_swa_retain_len),即位置 < L-retain+1。先回收再分配。 + 必须在 init_req_to_token_indexes 之后调用(位置对齐分配经 req_to_token 行派生/scatter)。""" + assert self.mem_manager is not None + if self.sliding_window is not None: + retain = self._swa_retain_len() + evict_slots = [] + req_list = b_req_idx.detach().cpu().tolist() + ready_list = b_ready_cache_len.detach().cpu().tolist() + for req_idx, ready_len in zip(req_list, ready_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + ready_len = int(ready_len) + mark = self._swa_evict_marks[req_idx] + if mark < 0: + # 首个 chunk: [0, ready_len) 是 radix 共享前缀,其 swa 槽归 radix 所有,不可回收。 + self._swa_evict_marks[req_idx] = ready_len + continue + evict_end = ready_len - retain + 1 + if evict_end > mark: + evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) + self._swa_evict_marks[req_idx] = evict_end + if evict_slots: + self.mem_manager.evict_swa(torch.cat(evict_slots)) + self.mem_manager.alloc_swa_prefill(b_req_idx, b_ready_cache_len, b_seq_len, self.req_to_token_indexs) return - def _reset_compress_cache_req(self, cache: LayerCache, req_idx: int): - if cache.layer_num == 0: - return - cache.buffer[:, req_idx, 0, ...].fill_(0) - cache.buffer[:, req_idx, 1, ...].fill_(float("-inf")) + def prepare_decode_swa( + self, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + mem_indexes: torch.Tensor, + ) -> None: + """decode prep: 回收出窗槽并为本步新 token 分配位置对齐的 swa 槽。当前 query 位置 + seq_len-1 的窗口是 [seq_len-W, seq_len-1];回收边界额外保留一个 radix 页 + (_swa_retain_len),即位置 < seq_len-retain。先回收再分配。""" + assert self.mem_manager is not None + if self.sliding_window is not None: + retain = self._swa_retain_len() + evict_slots = [] + req_list = b_req_idx.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + for req_idx, seq_len in zip(req_list, seq_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + seq_len = int(seq_len) + mark = self._swa_evict_marks[req_idx] + if mark < 0: + # 未经过 prefill prep 的保守路径: 不回收旧位置,仅推进水位线。 + self._swa_evict_marks[req_idx] = max(0, seq_len - retain) + continue + evict_end = seq_len - retain + if evict_end > mark: + evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) + self._swa_evict_marks[req_idx] = evict_end + if evict_slots: + self.mem_manager.evict_swa(torch.cat(evict_slots)) + self.mem_manager.alloc_swa_decode(b_req_idx, b_seq_len, mem_indexes, self.req_to_token_indexs) return def _reset_state_pool_req(self, cache: LayerCache, req_idx: int): @@ -465,105 +531,90 @@ def _reset_state_pool_req(self, cache: LayerCache, req_idx: int): return def init_compress_state(self, req_idx: int): - """新请求开始时重置其 compressor 在途状态(对应 mamba 的 init_linear_att_state)。""" + """新请求开始时重置其 compressor 在途状态(对应 mamba 的 init_linear_att_state)。 + + 只有 c128 状态是 req 键控的(c4 状态随 swa 页生灭,内核组起点覆写,无需重置; + 压缩槽位以 full 槽位为键,随请求 full 槽的释放级联回收)。""" self.clear_runtime_state(req_idx) - c4, c128 = self.pop_compress_indices_for_req(req_idx) - self.free_compress_indices(free_c4_index=c4, free_c128_index=c128) - if self.n_c4 > 0: - self._reset_compress_cache_req(self.req_to_c4_state, req_idx) - self._reset_compress_cache_req(self.req_to_c4_indexer_state, req_idx) - self._reset_state_pool_req(self.req_to_c4_state_pool, req_idx) - self._reset_state_pool_req(self.req_to_c4_indexer_state_pool, req_idx) if self.n_c128 > 0: - self._reset_compress_cache_req(self.req_to_c128_state, req_idx) self._reset_state_pool_req(self.req_to_c128_state_pool, req_idx) return - def _ensure_compress_slots(self, req_idx: int, ratio: int, entry_start: int, entry_count: int) -> torch.Tensor: - if entry_count == 0: - return torch.empty((0,), dtype=torch.int32, device="cuda") - assert entry_start >= 0 and entry_count >= 0 + # ------------------------------------------------------------------ compress slot prep (per step) + def _compress_mapping_alloc(self, ratio: int): assert self.mem_manager is not None, "DeepSeek-V4 mem manager is not bound yet" if ratio == 4: - table = self.req_to_c4_indexs - counts = self._c4_entry_counts - alloc = self.mem_manager.alloc_c4 - elif ratio == 128: - table = self.req_to_c128_indexs - counts = self._c128_entry_counts - alloc = self.mem_manager.alloc_c128 - else: - raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") - - required_count = entry_start + entry_count - assert required_count <= table.shape[1], ( - f"DeepSeek-V4 compressed slot table overflow: req={req_idx} " - f"ratio={ratio} required={required_count} capacity={table.shape[1]}" - ) - old_count = counts[req_idx] - if required_count > old_count: - new_slots_cpu = alloc(required_count - old_count) - table[req_idx, old_count:required_count] = new_slots_cpu.cuda(non_blocking=True) - counts[req_idx] = required_count - return table[req_idx, entry_start:required_count] - - def ensure_c4_slots(self, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: - return self._ensure_compress_slots(req_idx, 4, entry_start, entry_count) - - def ensure_c128_slots(self, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: - return self._ensure_compress_slots(req_idx, 128, entry_start, entry_count) - - def ensure_compress_slots(self, layer_index: int, req_idx: int, entry_start: int, entry_count: int) -> torch.Tensor: - ratio = self.compress_rates[layer_index] - if ratio == 4: - return self.ensure_c4_slots(req_idx, entry_start, entry_count) + return self.mem_manager.full_to_c4_indexs, self.mem_manager.alloc_c4 if ratio == 128: - return self.ensure_c128_slots(req_idx, entry_start, entry_count) - raise AssertionError(f"layer {layer_index} is not a compressed attention layer") + return self.mem_manager.full_to_c128_indexs, self.mem_manager.alloc_c128 + raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") - def prepare_decode_compress_slots(self, b_req_idx: torch.Tensor, b_seq_len: torch.Tensor) -> None: + def _scatter_compress_slots(self, ratio: int, full_slots: torch.Tensor) -> None: + """为组末 full 槽位分配压缩槽并写入映射。已映射(>=0)的行跳过——重复 prep 幂等。""" + if full_slots.numel() == 0: + return + mapping, alloc = self._compress_mapping_alloc(ratio) + full_slots = full_slots.cuda().long().reshape(-1) + # 去重: 同批重复键会让后写覆盖先写,先分配的压缩槽成为孤儿(allocator 泄漏)。 + need = torch.unique(full_slots[mapping[full_slots] < 0]) + if need.numel() == 0: + return + new_slots = alloc(need.numel()).cuda(non_blocking=True).to(torch.int32) + mapping[need] = new_slots + return + + def prepare_prefill_compress_slots( + self, + b_req_idx: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_seq_len: torch.Tensor, + ) -> None: + """prefill prep: 为本 chunk 内的组末 token(位置 (g+1)*ratio-1 ∈ [ready, seq))分配压缩槽, + scatter 进 full_to_c4/c128_indexs。必须在 init_req_to_token_indexes 之后(组末 full 槽 + 从 req_to_token_indexs 取)、attention metadata 构建之前调用。""" + if self.n_c4 == 0 and self.n_c128 == 0: + return req_list = b_req_idx.detach().cpu().tolist() + ready_list = b_ready_cache_len.detach().cpu().tolist() seq_list = b_seq_len.detach().cpu().tolist() - for req_idx, seq_len in zip(req_list, seq_list): - req_idx = int(req_idx) - if req_idx == self.HOLD_REQUEST_ID: + for ratio, n_layers in ((4, self.n_c4), (128, self.n_c128)): + if n_layers == 0: continue - seq_len = int(seq_len) - if self.n_c4 > 0: - required_c4 = seq_len // 4 - old_c4 = self._c4_entry_counts[req_idx] - if required_c4 > old_c4: - self.ensure_c4_slots(req_idx, old_c4, required_c4 - old_c4) - if self.n_c128 > 0: - required_c128 = seq_len // 128 - old_c128 = self._c128_entry_counts[req_idx] - if required_c128 > old_c128: - self.ensure_c128_slots(req_idx, old_c128, required_c128 - old_c128) + end_slots = [] + for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + first, last = int(ready_len) // ratio, int(seq_len) // ratio + if last > first: + ends = self.req_to_token_indexs[req_idx, ratio - 1 : last * ratio : ratio] + end_slots.append(ends[first:]) + if end_slots: + self._scatter_compress_slots(ratio, torch.cat(end_slots)) return - def pop_compress_indices_for_req(self, req_idx: int): - c4_count = self._c4_entry_counts[req_idx] - if c4_count > 0: - c4 = self.req_to_c4_indexs[req_idx, :c4_count].clone() - self.req_to_c4_indexs[req_idx, :c4_count].fill_(0) - self._c4_entry_counts[req_idx] = 0 - else: - c4 = None - - c128_count = self._c128_entry_counts[req_idx] - if c128_count > 0: - c128 = self.req_to_c128_indexs[req_idx, :c128_count].clone() - self.req_to_c128_indexs[req_idx, :c128_count].fill_(0) - self._c128_entry_counts[req_idx] = 0 - else: - c128 = None - return c4, c128 - - def free_compress_indices(self, free_c4_index=None, free_c128_index=None): - if free_c4_index is not None and len(free_c4_index) > 0: - self.mem_manager.free_c4(free_c4_index) - if free_c128_index is not None and len(free_c128_index) > 0: - self.mem_manager.free_c128(free_c128_index) + def prepare_decode_compress_slots( + self, + b_req_idx: torch.Tensor, + b_seq_len: torch.Tensor, + mem_indexes: torch.Tensor, + ) -> None: + """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 + 组末 full 槽即本步的 mem_index(此刻 req_to_token_indexs 尚未写入本步槽位)。""" + if self.n_c4 == 0 and self.n_c128 == 0: + return + req_list = b_req_idx.detach().cpu().tolist() + seq_list = b_seq_len.detach().cpu().tolist() + for ratio, n_layers in ((4, self.n_c4), (128, self.n_c128)): + if n_layers == 0: + continue + rows = [ + i + for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) + if int(req_idx) != self.HOLD_REQUEST_ID and int(seq_len) > 0 and int(seq_len) % ratio == 0 + ] + if rows: + self._scatter_compress_slots(ratio, mem_indexes.reshape(-1)[rows]) return def alloc(self): @@ -573,53 +624,17 @@ def alloc(self): return req_idx def clear_runtime_state(self, req_idx: int): - self._runtime_states[req_idx].clear() - if self.mem_manager is not None and hasattr(self.mem_manager, "free_swa_for_req"): - self.mem_manager.free_swa_for_req(req_idx) - return - - def set_runtime_state(self, req_idx: int, layer_index: int, state: dict): - self._runtime_states[req_idx][layer_index] = state + # swa 槽位本身由 mem_manager.free 级联回收(随 full 槽位),这里只复位出窗水位线。 + self._swa_evict_marks[req_idx] = -1 return - def get_runtime_state(self, req_idx: int, layer_index: int): - return self._runtime_states[req_idx][layer_index] - - def get_compress_state_for_req(self, layer_index: int, req_idx: int): - if self.compress_rates[layer_index] == 4: - state = self.get_c4_compress_state(layer_index) - elif self.compress_rates[layer_index] == 128: - state = self.get_c128_compress_state(layer_index) - else: - raise AssertionError(f"layer {layer_index} is not a compressed attention layer") - return state[req_idx, 0], state[req_idx, 1] - def get_compress_state_pool_for_req(self, layer_index: int, req_idx: int): - if self.compress_rates[layer_index] == 4: - cache = self.req_to_c4_state_pool - local = self.layer_to_c4_idx[layer_index] - elif self.compress_rates[layer_index] == 128: - cache = self.req_to_c128_state_pool - local = self.layer_to_c128_idx[layer_index] - else: - raise AssertionError(f"layer {layer_index} is not a compressed attention layer") - return cache.buffer[local, req_idx] - - def get_c4_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.layer_to_c4_idx[layer_index] - return self.req_to_c4_state.buffer[local] - - def get_c128_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.layer_to_c128_idx[layer_index] - return self.req_to_c128_state.buffer[local] + assert self.compress_rates[layer_index] == 128, "c4 state 在 mem manager 的 swa 页派生池" + return self.req_to_c128_state_pool.buffer[self.layer_to_c128_idx[layer_index], req_idx] - def get_c4_indexer_compress_state(self, layer_index: int) -> torch.Tensor: - local = self.layer_to_c4_idx[layer_index] - return self.req_to_c4_indexer_state.buffer[local] - - def get_c4_indexer_state_pool_for_req(self, layer_index: int, req_idx: int) -> torch.Tensor: - local = self.layer_to_c4_idx[layer_index] - return self.req_to_c4_indexer_state_pool.buffer[local, req_idx] + def get_compress_state_pool(self, layer_index: int): + assert self.compress_rates[layer_index] == 128, "c4 state 在 mem manager 的 swa 页派生池" + return self.req_to_c128_state_pool.buffer[self.layer_to_c128_idx[layer_index]] def get_prompt_cache_value_ops(self): return DeepseekV4PromptCacheValueOps(self) @@ -627,262 +642,76 @@ def get_prompt_cache_value_ops(self): def get_prompt_cache_page_size(self): return 128 - def _slice_cpu_slots(self, slots: Optional[torch.Tensor], start: int, end: int, ratio: int): - if slots is None: - return None - return slots[start // ratio : end // ratio].clone() - - def _slice_swa_payload(self, swa_payload, start: int, end: int): - if swa_payload is None: - return None - positions = swa_payload["positions"] - mask = (positions >= start) & (positions < end) - if not bool(mask.any()): - return None - return { - "positions": positions[mask].clone(), - "full_slots": swa_payload["full_slots"][mask].clone(), - "swa_slots": swa_payload["swa_slots"][mask].clone(), - } + def compute_swa_page_valid(self, full_slots: torch.Tensor) -> torch.Tensor: + """按当下 full_to_swa 映射给出按页有效性: full_slots [L](L 为 page 整数倍) -> + cpu bool [L/page],页内全部映射有效才为 True。GPU gather + 同步,测试/校验用; + 插入热路径用 swa_page_valid_from_watermark(纯 CPU,免同步)。""" + page = self.get_prompt_cache_page_size() + assert full_slots.numel() % page == 0 + if full_slots.numel() == 0: + return torch.zeros((0,), dtype=torch.bool) + swa = self.mem_manager.full_to_swa_indexs[full_slots.cuda().long().reshape(-1)] + return (swa.view(-1, page) >= 0).all(dim=1).cpu() + + def swa_page_valid_from_watermark(self, req_idx: int, cache_len: int) -> torch.Tensor: + """插入时的按页有效性,纯 CPU: 请求自有 token 的 swa 映射只被出窗水位线回收 + (阀不触活跃请求,级联只在 free 时),页 p 全驻留 ⟺ 页起点 128p >= 水位线。 + + 与 compute_swa_page_valid 在插入时刻对自有 token 等价,但不做 GPU gather/同步—— + router 关键路径上每次插入省一次对全部在途 kernel 的等待。bitmap 中借入前缀 + ([0, ready) 的页)的行在 radix insert 切片时被丢弃(既有节点保留自己的 bitmap), + 其取值无影响。""" + page = self.get_prompt_cache_page_size() + mark = max(0, self._swa_evict_marks[req_idx]) + n_pages = int(cache_len) // page + return torch.arange(n_pages, dtype=torch.long) * page >= mark def slice_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload, start: int, end: int): start = int(start) end = int(end) - # c4/c128/indexer-K slots are true historical KV and can be sliced by ratio. - # compressor running state only describes the payload end boundary; it is valid - # for a slice only when that slice keeps the original end boundary. - keep_end_state = end == payload.cache_len + page = self.get_prompt_cache_page_size() + # radix page=128 保证分裂点页对齐,bitmap 可整页切分。 return DeepseekV4PromptCachePayload( cache_len=end - start, - c4_slots=self._slice_cpu_slots(payload.c4_slots, start, end, 4), - c128_slots=self._slice_cpu_slots(payload.c128_slots, start, end, 128), - c4_state=payload.c4_state.clone() if keep_end_state and payload.c4_state is not None else None, - c4_state_pool=payload.c4_state_pool.clone() - if keep_end_state and payload.c4_state_pool is not None - else None, - c4_indexer_state=payload.c4_indexer_state.clone() - if keep_end_state and payload.c4_indexer_state is not None - else None, - c4_indexer_state_pool=payload.c4_indexer_state_pool.clone() - if keep_end_state and payload.c4_indexer_state_pool is not None + swa_page_valid=payload.swa_page_valid[start // page : end // page].clone() + if payload.swa_page_valid is not None else None, - swa=self._slice_swa_payload(payload.swa, start, end), ) def concat_prompt_cache_payloads(self, payloads: List[DeepseekV4PromptCachePayload]): if len(payloads) == 0: return None - c4_slots = [p.c4_slots for p in payloads if p.c4_slots is not None and len(p.c4_slots) > 0] - c128_slots = [p.c128_slots for p in payloads if p.c128_slots is not None and len(p.c128_slots) > 0] - last = payloads[-1] + bitmaps = [p.swa_page_valid for p in payloads] return DeepseekV4PromptCachePayload( cache_len=sum(p.cache_len for p in payloads), - c4_slots=torch.cat(c4_slots, dim=0) if c4_slots else None, - c128_slots=torch.cat(c128_slots, dim=0) if c128_slots else None, - c4_state=last.c4_state, - c4_state_pool=last.c4_state_pool, - c4_indexer_state=last.c4_indexer_state, - c4_indexer_state_pool=last.c4_indexer_state_pool, - swa=last.swa, + swa_page_valid=torch.cat(bitmaps, dim=0) if all(b is not None for b in bitmaps) else None, ) def build_prompt_cache_payload( self, req_idx: int, cache_len: int, - clone_swa: bool = False, ) -> DeepseekV4PromptCachePayload: + """构造插入载荷。compressor 状态不进载荷(c4 随 swa 页生灭、c128 边界自然归零), + cache_len 不再受序列末端约束——任意 128 对齐前缀皆可插入。 + swa_page_valid 不在此填: 它必须用插入时刻的映射(infer batch 在 insert 前补)。""" assert self.mem_manager is not None - cache_len = int(cache_len) - full_slots = self.req_to_token_indexs[req_idx, :cache_len].detach().cpu() - c4_count = cache_len // 4 - c128_count = cache_len // 128 - c4_slots = self.req_to_c4_indexs[req_idx, :c4_count].detach().cpu().clone() if c4_count > 0 else None - c128_slots = self.req_to_c128_indexs[req_idx, :c128_count].detach().cpu().clone() if c128_count > 0 else None - if clone_swa: - swa_payload = self.mem_manager.clone_swa_for_prompt_cache(req_idx, cache_len, full_slots) - else: - swa_payload = self.mem_manager.snapshot_swa_for_prompt_cache(req_idx, cache_len, full_slots) - return DeepseekV4PromptCachePayload( - cache_len=cache_len, - c4_slots=c4_slots, - c128_slots=c128_slots, - c4_state=self.req_to_c4_state.buffer[:, req_idx].detach().clone() if self.n_c4 > 0 else None, - c4_state_pool=self.req_to_c4_state_pool.buffer[:, req_idx].detach().clone() if self.n_c4 > 0 else None, - c4_indexer_state=self.req_to_c4_indexer_state.buffer[:, req_idx].detach().clone() - if self.n_c4 > 0 - else None, - c4_indexer_state_pool=self.req_to_c4_indexer_state_pool.buffer[:, req_idx].detach().clone() - if self.n_c4 > 0 - else None, - swa=swa_payload, - ) - - def detach_prompt_cache_payload_from_req(self, req_idx: int, payload: DeepseekV4PromptCachePayload): - if payload is not None and self.mem_manager is not None: - self.mem_manager.detach_swa_for_prompt_cache(req_idx, payload.swa) - return - - def free_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload): - if payload is None or self.mem_manager is None: - return - if payload.c4_slots is not None and len(payload.c4_slots) > 0: - self.mem_manager.free_c4(payload.c4_slots) - if payload.c128_slots is not None and len(payload.c128_slots) > 0: - self.mem_manager.free_c128(payload.c128_slots) - self.mem_manager.free_swa_prompt_cache(payload.swa) - return - - def release_prompt_cache_detached_swa( - self, - payload: DeepseekV4PromptCachePayload, - keep_payload: Optional[DeepseekV4PromptCachePayload] = None, - ): - if payload is None or payload.swa is None or self.mem_manager is None: - return - old_swa = payload.swa - if keep_payload is None or keep_payload.swa is None: - self.mem_manager.free_swa_prompt_cache(old_swa) - return - - old_slots = old_swa["swa_slots"].long() - keep_slots = keep_payload.swa["swa_slots"].long() - if old_slots.numel() == 0: - return - if keep_slots.numel() == 0: - self.mem_manager.free_swa_prompt_cache(old_swa) - return - - release_mask = ~torch.isin(old_slots, keep_slots) - if not release_mask.any(): - return - release_payload = { - "full_slots": old_swa["full_slots"][release_mask].clone(), - "swa_slots": old_swa["swa_slots"][release_mask].clone(), - } - self.mem_manager.free_swa_prompt_cache(release_payload) - return - - def _reset_c128_for_prompt_cache(self, req_idx: int): - if self.n_c128 > 0: - self._reset_compress_cache_req(self.req_to_c128_state, req_idx) - self._reset_state_pool_req(self.req_to_c128_state_pool, req_idx) - return - - def rebuild_runtime_state_for_req(self, req_idx: int): - state_map = self._runtime_states[req_idx] - state_map.clear() - for layer_index, ratio in enumerate(self.compress_rates): - if ratio == 4: - cstate_kv, cstate_score = self.get_compress_state_for_req(layer_index, req_idx) - idx_state = self.get_c4_indexer_compress_state(layer_index) - state_map[layer_index] = { - "cstate_kv": cstate_kv, - "cstate_score": cstate_score, - "idx_cstate_kv": idx_state[req_idx, 0], - "idx_cstate_score": idx_state[req_idx, 1], - } - elif ratio == 128: - cstate_kv, cstate_score = self.get_compress_state_for_req(layer_index, req_idx) - state_map[layer_index] = { - "cstate_kv": cstate_kv, - "cstate_score": cstate_score, - } - return - - def restore_prompt_cache_payload(self, req_idx: int, payload: DeepseekV4PromptCachePayload): - assert self.mem_manager is not None - cache_len = int(payload.cache_len) - c4_count = cache_len // 4 - c128_count = cache_len // 128 - if c4_count > 0: - assert payload.c4_slots is not None and len(payload.c4_slots) == c4_count - self.req_to_c4_indexs[req_idx, :c4_count] = payload.c4_slots.cuda(non_blocking=True) - if c128_count > 0: - assert payload.c128_slots is not None and len(payload.c128_slots) == c128_count - self.req_to_c128_indexs[req_idx, :c128_count] = payload.c128_slots.cuda(non_blocking=True) - self._c4_entry_counts[req_idx] = c4_count - self._c128_entry_counts[req_idx] = c128_count - - if self.n_c4 > 0: - if payload.c4_state is None or payload.c4_indexer_state is None: - raise RuntimeError("DeepSeek-V4 prompt cache hit is missing c4 running state") - self.req_to_c4_state.buffer[:, req_idx].copy_(payload.c4_state) - self.req_to_c4_indexer_state.buffer[:, req_idx].copy_(payload.c4_indexer_state) - if payload.c4_state_pool is not None: - self.req_to_c4_state_pool.buffer[:, req_idx].copy_(payload.c4_state_pool) - if payload.c4_indexer_state_pool is not None: - self.req_to_c4_indexer_state_pool.buffer[:, req_idx].copy_(payload.c4_indexer_state_pool) - self._reset_c128_for_prompt_cache(req_idx) - self.mem_manager.restore_swa_from_prompt_cache(payload.swa) - self.rebuild_runtime_state_for_req(req_idx) - return + return DeepseekV4PromptCachePayload(cache_len=int(cache_len)) - def pop_prompt_cache_free_compress_indices( - self, - req_idx: int, - keep_len: int, - duplicate_start_len: Optional[int] = None, - duplicate_end_len: Optional[int] = None, - ): - def collect(table, cur_count, ratio): - ranges = [] - if duplicate_start_len is not None and duplicate_end_len is not None: - dup_start = duplicate_start_len // ratio - dup_end = duplicate_end_len // ratio - if dup_end > dup_start: - ranges.append((dup_start, dup_end)) - keep_count = keep_len // ratio - if cur_count > keep_count: - ranges.append((keep_count, cur_count)) - parts = [table[req_idx, s:e].clone() for s, e in ranges if e > s] - return torch.cat(parts, dim=0) if parts else None - - c4 = collect(self.req_to_c4_indexs, self._c4_entry_counts[req_idx], 4) - c128 = collect(self.req_to_c128_indexs, self._c128_entry_counts[req_idx], 128) - if self._c4_entry_counts[req_idx] > 0: - self.req_to_c4_indexs[req_idx, : self._c4_entry_counts[req_idx]].fill_(0) - if self._c128_entry_counts[req_idx] > 0: - self.req_to_c128_indexs[req_idx, : self._c128_entry_counts[req_idx]].fill_(0) - self._c4_entry_counts[req_idx] = 0 - self._c128_entry_counts[req_idx] = 0 - return c4, c128 - - def free( - self, - free_req_indexes, - free_token_index, - free_c4_index=None, - free_c128_index=None, - ): - """释放 dense 槽(基类)+ 压缩槽。压缩槽由调用方(infer batch)从 req_to_c*_indexs 收集后传入, - 与基类用 free_token_index 传 dense 槽的方式一致。""" + def free(self, free_req_indexes, free_token_index): + """dense/swa/压缩槽全部经 mem_manager.free(free_token_index) 级联回收。""" for req_index in free_req_indexes: self.clear_runtime_state(req_index) super().free(free_req_indexes, free_token_index) - self.free_compress_indices(free_c4_index=free_c4_index, free_c128_index=free_c128_index) return def free_req(self, free_req_index: int): self.clear_runtime_state(free_req_index) - c4, c128 = self.pop_compress_indices_for_req(free_req_index) - self.free_compress_indices(free_c4_index=c4, free_c128_index=c128) return super().free_req(free_req_index) def free_all(self): super().free_all() - self._runtime_states = [{} for _ in range(self.max_request_num + 1)] - self._c4_entry_counts = [0 for _ in range(self.max_request_num + 1)] - self._c128_entry_counts = [0 for _ in range(self.max_request_num + 1)] - if self.n_c4 > 0: - self.req_to_c4_indexs.fill_(0) - self.req_to_c4_state.buffer.fill_(0) - self.req_to_c4_indexer_state.buffer.fill_(0) - self.req_to_c4_state_pool.buffer.fill_(0) - self.req_to_c4_indexer_state_pool.buffer.fill_(0) + self._swa_evict_marks = [-1 for _ in range(self.max_request_num + 1)] if self.n_c128 > 0: - self.req_to_c128_indexs.fill_(0) - self.req_to_c128_state.buffer.fill_(0) self.req_to_c128_state_pool.buffer.fill_(0) - self._init_all_score_state() return diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index d0c2745161..39a6889d72 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -9,8 +9,8 @@ class DeepseekV4InferStateInfo(InferStateInfo): mem_manager: DeepseekV4MemoryManager """Per-token interleaved-rope cos/sin for the two rope variants (sliding / compressed), following - the gemma4 two-variant convention (_cos_cached_* -> position_cos_*). Also exposes the full compressed - cos/sin tables, which the KV compressor indexes at window positions (not per-token).""" + the gemma4 two-variant convention (_cos_cached_* -> position_cos_*). The full rope tables are + model constants and live on the model / layer infers, not here.""" def __init__(self): super().__init__() @@ -18,8 +18,6 @@ def __init__(self): self.position_sin_sliding = None self.position_cos_compress = None self.position_sin_compress = None - self.cos_compress_table = None - self.sin_compress_table = None def init_some_extra_state(self, model): super().init_some_extra_state(model) # sets position_ids, b_q_seq_len, b_q_start_loc (prefill) @@ -28,5 +26,3 @@ def init_some_extra_state(self, model): self.position_sin_sliding = torch.index_select(model._sin_cached_sliding, 0, pos) self.position_cos_compress = torch.index_select(model._cos_cached_compress, 0, pos) self.position_sin_compress = torch.index_select(model._sin_cached_compress, 0, pos) - self.cos_compress_table = model._cos_cached_compress - self.sin_compress_table = model._sin_cached_compress diff --git a/lightllm/models/deepseek_v4/layer_infer/attention.py b/lightllm/models/deepseek_v4/layer_infer/attention.py deleted file mode 100644 index 8a7428f0dd..0000000000 --- a/lightllm/models/deepseek_v4/layer_infer/attention.py +++ /dev/null @@ -1,149 +0,0 @@ -import os - -import torch - - -FLASHMLA_MIN_HEADS = 64 -FLASHMLA_TOPK_MULTIPLE = 128 -DSV4_DEBUG_TORCH_SPARSE_ATTN = os.getenv("DSV4_DEBUG_TORCH_SPARSE_ATTN", "0") == "1" - - -def _pad_topk_for_flashmla(topk_idxs): - K = topk_idxs.shape[-1] - padded_K = ((K + FLASHMLA_TOPK_MULTIPLE - 1) // FLASHMLA_TOPK_MULTIPLE) * FLASHMLA_TOPK_MULTIPLE - if padded_K == K: - return topk_idxs.contiguous() - padded = torch.full((*topk_idxs.shape[:-1], padded_K), -1, device=topk_idxs.device, dtype=topk_idxs.dtype) - padded[..., :K] = topk_idxs - return padded.contiguous() - - -def _compact_topk_indices(topk_idxs, kv_len): - valid = (topk_idxs >= 0) & (topk_idxs < kv_len) - topk_lens = valid.sum(dim=-1).to(torch.int32) - if valid.all(): - return topk_idxs.contiguous(), topk_lens.contiguous() - - compact = torch.full_like(topk_idxs, -1) - ranks = valid.to(torch.int32).cumsum(dim=-1) - 1 - rows = torch.arange(topk_idxs.shape[0], device=topk_idxs.device).unsqueeze(1).expand_as(topk_idxs) - compact[rows[valid], ranks[valid].long()] = topk_idxs[valid] - return compact.contiguous(), topk_lens.contiguous() - - -def _pad_heads_for_flashmla(q, attn_sink): - h = q.shape[1] - if h == FLASHMLA_MIN_HEADS: - return q.contiguous(), attn_sink.to(torch.float32).contiguous(), h - if h > FLASHMLA_MIN_HEADS: - raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention only supports up to 64 local heads, got {h}") - - q_pad = q.new_zeros(q.shape[0], FLASHMLA_MIN_HEADS, q.shape[2]) - q_pad[:, :h] = q - sink_pad = torch.full((FLASHMLA_MIN_HEADS,), -float("inf"), device=q.device, dtype=torch.float32) - sink_pad[:h] = attn_sink.to(torch.float32) - return q_pad.contiguous(), sink_pad.contiguous(), h - - -def _torch_sparse_attn(q, kv, attn_sink, topk_idxs, scale): - return _torch_sparse_attn_flat(q[0], kv[0], attn_sink, topk_idxs[0], scale).unsqueeze(0) - - -def _torch_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale): - q0 = q.float() - kv0 = kv.float() - indices = topk_idxs.long() - valid = (indices >= 0) & (indices < kv0.shape[0]) - safe_indices = torch.where(valid, indices, torch.zeros_like(indices)) - kv_sel = kv0[safe_indices] - scores = torch.einsum("mhd,mkd->mhk", q0, kv_sel) * scale - scores = scores.masked_fill(~valid.unsqueeze(1), float("-inf")) - sink = attn_sink.float().view(1, -1) - max_scores = torch.maximum(scores.max(dim=-1).values, sink) - exp_scores = torch.exp(scores - max_scores.unsqueeze(-1)).masked_fill(~valid.unsqueeze(1), 0.0) - exp_sink = torch.exp(sink - max_scores) - denom = exp_scores.sum(dim=-1) + exp_sink - out = torch.einsum("mhk,mkd->mhd", exp_scores / denom.unsqueeze(-1), kv_sel) - return out.to(q.dtype) - - -def vllm_sparse_attn(q, kv, attn_sink, topk_idxs, scale): - """DeepSeek-V4 sparse MLA through vLLM FlashMLA. - - q:[1,m,h,d], kv:[1,n,d] (single KV head shared over h), attn_sink:[h], - topk_idxs:[1,m,K] int (-1 = invalid/skip). Returns o:[1,m,h,d]. - """ - b, m, h, d = q.shape - if b != 1 or kv.shape[0] != 1 or topk_idxs.shape[0] != 1: - raise RuntimeError("DeepSeek-V4 FlashMLA sparse attention wrapper expects one request per call") - if d != 512: - raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires head_dim=512, got {d}") - if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: - raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires bf16 q/kv, got {q.dtype}/{kv.dtype}") - - return vllm_sparse_attn_flat(q[0], kv[0], attn_sink, topk_idxs[0], scale).unsqueeze(0) - - -def vllm_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale, already_compact=False): - """FlashMLA sparse attention over a flat KV arena. - - q:[m,h,d], kv:[n,d], topk_idxs:[m,K] int. Indices are global offsets into - the flat kv tensor, so callers can concatenate per-request KV candidates and - run one FlashMLA call for the whole batch. When already_compact=True, each - row must place all valid indices before invalid (-1) entries. - """ - m, h, d = q.shape - if d != 512: - raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires head_dim=512, got {d}") - if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: - raise RuntimeError(f"DeepSeek-V4 FlashMLA sparse attention requires bf16 q/kv, got {q.dtype}/{kv.dtype}") - if q.shape[0] == 0: - return q.new_empty((0, h, d)) - - if DSV4_DEBUG_TORCH_SPARSE_ATTN: - return _torch_sparse_attn_flat(q, kv, attn_sink, topk_idxs, scale) - - from vllm.third_party.flashmla.flash_mla_interface import flash_mla_sparse_fwd - - q_pad, sink_pad, real_heads = _pad_heads_for_flashmla(q, attn_sink) - topk_idxs = topk_idxs.to(torch.int32) - if already_compact: - valid = (topk_idxs >= 0) & (topk_idxs < kv.shape[0]) - indices = topk_idxs.contiguous() - topk_lens = valid.sum(dim=-1).to(torch.int32).contiguous() - else: - indices, topk_lens = _compact_topk_indices(topk_idxs, kv.shape[0]) - indices = _pad_topk_for_flashmla(indices).unsqueeze(1) - kv_flat = kv.unsqueeze(1).contiguous() - out, _, _ = flash_mla_sparse_fwd( - q=q_pad, - kv=kv_flat, - indices=indices, - sm_scale=scale, - attn_sink=sink_pad, - topk_length=topk_lens, - out=None, - ) - return out[:, :real_heads].to(q.dtype) - - -def build_prefill_topk_idxs(seqlen, window, ratio, n_window, device): - """Per-query candidate indices into [window_kv (n_window tokens) ++ compressed_kv (ncomp entries)]. - - Returns int32 [seqlen, window + ncomp] with -1 for invalid. Window part indexes the per-token KV - (here stored as tokens 0..seqlen-1, so n_window == seqlen); compressed part is offset by n_window. - For prompts where ncomp <= index_topk the indexer is a no-op, so all causally-valid compressed - entries are attended (matches the reference for short context). - """ - t = torch.arange(seqlen, device=device) - offsets = torch.arange(window, device=device) - win = t.unsqueeze(1) - (window - 1 - offsets).unsqueeze(0) - win = torch.where(win >= 0, win, torch.full_like(win, -1)) - if ratio: - ncomp = seqlen // ratio - c = torch.arange(ncomp, device=device) - comp_valid = c.unsqueeze(0) < ((t.unsqueeze(1) + 1) // ratio) # [s, ncomp] - comp_idx = (c.unsqueeze(0) + n_window).expand(seqlen, ncomp) - comp = torch.where(comp_valid, comp_idx, torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long)) - return torch.cat([win, comp], dim=1).int() - return win.int() diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index f51f73829c..2256ecd1a9 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -1,121 +1,37 @@ -import importlib.util -import logging -import sys -import types -from pathlib import Path - import torch -import torch.nn.functional as F -from ..triton_kernel.rotary_emb import apply_rotary_emb -logger = logging.getLogger(__name__) -_SGLANG_COMPRESS_MOD = None _SGLANG_COMPRESS_ERR = None -_SGLANG_COMPRESS_WARNED = False +_SGLANG_COMPRESS_MOD = None +_SGLANG_LINEAR_BF16_FP32 = None _FREQ_CIS_CACHE = {} -# KV compressor: pools every `ratio` consecutive tokens into one compressed KV entry via gated -# (softmax) pooling + a learned absolute-position bias (ape), RMSNorm, and rope on the trailing -# rope_dim. ratio==4 uses overlapping windows (two-series Ca/Cb scheme). Pure-torch transcription of -# the bundled reference inference/model.py Compressor.forward for the prefill (start_pos==0) path. -# NOTE: the reference also applies an FP8/FP4 QAT activation sim to the compressed entry; omitted here -# for the correctness-first prefill path (negligible vs argmax; revisit if e2e diverges). - - -def _overlap_transform(tensor, ratio, d, value): - # tensor: [nwin, ratio, 2*d] -> [nwin, 2*ratio, d]; slots [ratio:]=Cb(current), [:ratio]=Ca(previous window) - nwin = tensor.shape[0] - out = tensor.new_full((nwin, 2 * ratio, d), value) - out[:, ratio:] = tensor[:, :, d:] - out[1:, :ratio] = tensor[:-1, :, :d] - return out - - -def _rmsnorm(x, weight, eps): - xf = x.float() - xf = xf * torch.rsqrt(xf.square().mean(-1, keepdim=True) + eps) - return (xf * weight.float()).to(x.dtype) - - -def _load_file_module(name, path): - spec = importlib.util.spec_from_file_location(name, path) - mod = importlib.util.module_from_spec(spec) - sys.modules[name] = mod - spec.loader.exec_module(mod) - return mod - def _load_sglang_compressor(): - global _SGLANG_COMPRESS_MOD, _SGLANG_COMPRESS_ERR + global _SGLANG_COMPRESS_ERR, _SGLANG_COMPRESS_MOD, _SGLANG_LINEAR_BF16_FP32 if _SGLANG_COMPRESS_MOD is not None: - return _SGLANG_COMPRESS_MOD + return _SGLANG_COMPRESS_MOD, _SGLANG_LINEAR_BF16_FP32 if _SGLANG_COMPRESS_ERR is not None: raise _SGLANG_COMPRESS_ERR try: - from sglang.jit_kernel.dsv4 import compress_old as mod - - _SGLANG_COMPRESS_MOD = mod - return mod - except Exception as first_exc: - root = Path("/data/wanzihao/sglang/python/sglang") - try: - if not root.exists(): - raise first_exc - if "sglang" not in sys.modules: - sglang_mod = types.ModuleType("sglang") - sglang_mod.__path__ = [str(root)] - sys.modules["sglang"] = sglang_mod - if "sglang.utils" not in sys.modules: - utils_mod = types.ModuleType("sglang.utils") - utils_mod.is_in_ci = lambda: False - sys.modules["sglang.utils"] = utils_mod - if "sglang.jit_kernel" not in sys.modules: - jit_mod = types.ModuleType("sglang.jit_kernel") - jit_mod.__path__ = [str(root / "jit_kernel")] - sys.modules["sglang.jit_kernel"] = jit_mod - if "sglang.jit_kernel.dsv4" not in sys.modules: - dsv4_mod = types.ModuleType("sglang.jit_kernel.dsv4") - dsv4_mod.__path__ = [str(root / "jit_kernel" / "dsv4")] - sys.modules["sglang.jit_kernel.dsv4"] = dsv4_mod - if "sglang.srt" not in sys.modules: - srt_mod = types.ModuleType("sglang.srt") - srt_mod.__path__ = [str(root / "srt")] - sys.modules["sglang.srt"] = srt_mod - if "sglang.srt.environ" not in sys.modules: - env_mod = types.ModuleType("sglang.srt.environ") - - class _FalseEnv: - def get(self): - return False - - class _Envs: - SGLANG_OPT_USE_ONLINE_COMPRESS = _FalseEnv() - - env_mod.envs = _Envs() - sys.modules["sglang.srt.environ"] = env_mod - if "sglang.jit_kernel.utils" not in sys.modules: - _load_file_module("sglang.jit_kernel.utils", root / "jit_kernel" / "utils.py") - if "sglang.jit_kernel.dsv4.utils" not in sys.modules: - _load_file_module( - "sglang.jit_kernel.dsv4.utils", - root / "jit_kernel" / "dsv4" / "utils.py", - ) - _SGLANG_COMPRESS_MOD = _load_file_module( - "sglang.jit_kernel.dsv4.compress_old", - root / "jit_kernel" / "dsv4" / "compress_old.py", - ) - return _SGLANG_COMPRESS_MOD - except Exception as exc: - _SGLANG_COMPRESS_ERR = exc - raise exc - - -def _warn_sglang_fallback(exc): - global _SGLANG_COMPRESS_WARNED - if not _SGLANG_COMPRESS_WARNED: - logger.warning("DeepSeek-V4 SGLang compressor JIT unavailable, fallback to torch: %s", exc) - _SGLANG_COMPRESS_WARNED = True + from sglang.jit_kernel.dsv4 import linear_bf16_fp32 + from sglang.jit_kernel.dsv4 import compress_old as compress_mod + except Exception as exc: + _SGLANG_COMPRESS_ERR = RuntimeError( + "DeepSeek-V4 fused compressor requires sglang.jit_kernel.dsv4 " + "(linear_bf16_fp32 + compress_old). Install/export the SGLang package " + "or vendor the DSv4 compressor JIT into LightLLM." + ) + raise _SGLANG_COMPRESS_ERR from exc + _SGLANG_COMPRESS_MOD = compress_mod + _SGLANG_LINEAR_BF16_FP32 = linear_bf16_fp32 + return compress_mod, linear_bf16_fp32 + + +def _load_paged_compress_data_fn(): + from sglang.jit_kernel.dsv4 import triton_create_paged_compress_data + + return triton_create_paged_compress_data def _freq_cis(cos_table, sin_table): @@ -139,39 +55,27 @@ def _sglang_ape(ape, ratio, head_dim): return ape.contiguous() -def _pack_kv_score(kv, score, ratio, head_dim): - if ratio == 4: - return torch.cat( - [ - kv[:, :head_dim], - kv[:, head_dim:], - score[:, :head_dim], - score[:, head_dim:], - ], - dim=1, - ).contiguous() - return torch.cat([kv, score], dim=1).contiguous() - - -def _build_state_from_kv_score(kv, score, ape, ratio, head_dim): - overlap = ratio == 4 - kv_state, score_state = new_compressor_state(ratio, head_dim, kv.device) - s = kv.shape[0] - remainder = s % ratio - cutoff = s - remainder - offset = ratio if overlap else 0 - if overlap and cutoff >= ratio: - kv_state[:ratio] = kv[cutoff - ratio : cutoff] - score_state[:ratio] = score[cutoff - ratio : cutoff] + ape.float() - if remainder > 0: - kv_state[offset : offset + remainder] = kv[cutoff:] - score_state[offset : offset + remainder] = score[cutoff:] + ape.float()[:remainder] - return kv_state, score_state - - -def _sglang_prefill_from_kv_score( - kv, - score, +def _compressor_weight(wkv_w, wgate_w): + return torch.cat([wkv_w, wgate_w], dim=0).contiguous() + + +def _project_kv_score(x, wkv_w, wgate_w): + _, linear_bf16_fp32 = _load_sglang_compressor() + return linear_bf16_fp32(x, _compressor_weight(wkv_w, wgate_w)) + + +def _state_pool_view(state_pool): + if state_pool is None: + raise RuntimeError("DeepSeek-V4 fused compressor requires a persistent state_pool") + if state_pool.dim() == 4 and state_pool.shape[1] == 1: + return state_pool.squeeze(1) + return state_pool + + +def compressor_prefill_state( + x, + wkv_w, + wgate_w, norm_w, ape, ratio, @@ -179,48 +83,42 @@ def _sglang_prefill_from_kv_score( cos_table, sin_table, eps, - dtype, - state_pool=None, + state_pool, ): - if not kv.is_cuda or head_dim % 128 != 0 or ratio not in (4, 128): - return None, None - mod = _load_sglang_compressor() - kv_score = _pack_kv_score(kv, score, ratio, head_dim) - ape_sglang = _sglang_ape(ape.float(), ratio, head_dim) - slots = 8 if ratio == 4 else ratio - if state_pool is None: - state_pool = torch.zeros((1, slots, kv_score.shape[1]), device=kv.device, dtype=kv_score.dtype) - else: - state_pool.zero_() - seq_len = kv.shape[0] + """start_pos==0 prefill for ONE request: x [s, dim] -> compressed entries [s//ratio, head_dim] + (rope applied). state_pool is the request's persistent jit state slice [1, slots, coff*2*head_dim]; + it is rebuilt in place so the decode path can continue from the trailing partial window.""" + mod, _ = _load_sglang_compressor() + kv_score = _project_kv_score(x, wkv_w, wgate_w) + pool = _state_pool_view(state_pool) + pool.zero_() + seq_len = x.shape[0] plan = mod.CompressorPrefillPlan.generate( ratio, seq_len, torch.tensor([seq_len], dtype=torch.int64), torch.tensor([seq_len], dtype=torch.int64), - kv.device, + x.device, ) - indices = torch.zeros((1,), device=kv.device, dtype=torch.int32) + indices = torch.zeros((1,), device=x.device, dtype=torch.int32) out = mod.compress_forward( - state_pool, + pool, kv_score, - ape_sglang, + _sglang_ape(ape.float(), ratio, head_dim), indices, plan, head_dim=head_dim, compress_ratio=ratio, ) ncomp = seq_len // ratio - if ncomp: - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() - out = out.index_select(0, ragged_ids).to(dtype) - else: - out = kv.new_zeros(0, head_dim).to(dtype) - return out, state_pool + if ncomp == 0: + return x.new_zeros(0, head_dim) + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() + return out.index_select(0, ragged_ids).to(x.dtype) -def _sglang_decode_step_from_state_pool( +def compressor_decode_step_single( x_new, wkv_w, wgate_w, @@ -231,17 +129,14 @@ def _sglang_decode_step_from_state_pool( cos_table, sin_table, eps, - start_pos, state_pool, + start_pos, ): - if state_pool is None or not x_new.is_cuda or head_dim % 128 != 0 or ratio not in (4, 128): - return None, False - mod = _load_sglang_compressor() - xf = x_new.float().view(1, -1) - kv = F.linear(xf, wkv_w.float()) - score = F.linear(xf, wgate_w.float()) - kv_score = _pack_kv_score(kv, score, ratio, head_dim) - ape_sglang = _sglang_ape(ape.float(), ratio, head_dim) + """One token for ONE request (chunked-prefill extend path). Returns the finished compressed + entry [head_dim] when (start_pos+1) % ratio == 0, else None. Mutates state_pool in place.""" + mod, _ = _load_sglang_compressor() + kv_score = _project_kv_score(x_new.view(1, -1), wkv_w, wgate_w) + pool = _state_pool_view(state_pool) seq_len = start_pos + 1 plan = mod.CompressorDecodePlan( ratio, @@ -249,66 +144,22 @@ def _sglang_decode_step_from_state_pool( ) indices = torch.zeros((1,), device=x_new.device, dtype=torch.int32) out = mod.compress_forward( - state_pool, + pool, kv_score, - ape_sglang, + _sglang_ape(ape.float(), ratio, head_dim), indices, plan, head_dim=head_dim, compress_ratio=ratio, ) if seq_len % ratio != 0: - return None, True + return None mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - return out[0].to(x_new.dtype), True - - -def compress_prefill(x, wkv_w, wgate_w, norm_w, ape, ratio, head_dim, rope_dim, cos_table, sin_table, eps): - """x:[s,dim] (one request, start_pos=0) -> compressed kv [nwin, head_dim] (rope applied to last rope_dim). - - nwin = s // ratio (remainder tokens are decode-state, handled in the decode path). wkv_w/wgate_w: - [coff*head_dim, dim]; norm_w:[head_dim]; ape:[ratio, coff*head_dim]; cos_table/sin_table: compress rope tables. - """ - overlap = ratio == 4 - coff = 2 if overlap else 1 - d = head_dim - s = x.shape[0] - nwin = s // ratio - if nwin == 0: - # fewer than `ratio` tokens -> no completed window -> no compressed entry (matches reference) - return x.new_zeros(0, head_dim) - cutoff = nwin * ratio - xf = x.float() - kv = F.linear(xf, wkv_w.float())[:cutoff].view(nwin, ratio, coff * d) - score = F.linear(xf, wgate_w.float())[:cutoff].view(nwin, ratio, coff * d) + ape.float() - if overlap: - kv = _overlap_transform(kv, ratio, d, 0.0) - score = _overlap_transform(score, ratio, d, float("-inf")) - kv = (kv * torch.softmax(score, dim=1)).sum(dim=1) # [nwin, d] fp32 - kv = _rmsnorm(kv.to(x.dtype), norm_w, eps) # [nwin, d] - pos = torch.arange(nwin, device=x.device) * ratio - kv_rope = apply_rotary_emb(kv[:, -rope_dim:], cos_table[pos], sin_table[pos]) # cos/sin: [nwin, rope_dim//2] - return torch.cat([kv[:, :-rope_dim], kv_rope], dim=1) - - -def new_compressor_state(ratio, head_dim, device, dtype=torch.float32): - """Per-request compressor running state (matches reference Compressor.kv_state/score_state).""" - coff = 2 if ratio == 4 else 1 - kv_state = torch.zeros(coff * ratio, coff * head_dim, device=device, dtype=dtype) - score_state = torch.full((coff * ratio, coff * head_dim), float("-inf"), device=device, dtype=dtype) - return kv_state, score_state - - -def _finish_entry(kv, norm_w, ape_unused, rope_dim, cos_table, sin_table, position, eps, dtype): - kv = _rmsnorm(kv.to(dtype), norm_w, eps) # [d] - cos = cos_table[position : position + 1] # [1, rope_dim//2] - sin = sin_table[position : position + 1] - kv_rope = apply_rotary_emb(kv[-rope_dim:].unsqueeze(0), cos, sin)[0] - return torch.cat([kv[:-rope_dim], kv_rope], dim=0) + return out[0].to(x_new.dtype) -def compressor_prefill_state( - x, +def compressor_decode_step_batch( + x_new, wkv_w, wgate_w, norm_w, @@ -319,207 +170,193 @@ def compressor_prefill_state( cos_table, sin_table, eps, - return_state_pool=False, - state_pool=None, + state_pool, + b_req_idx, + start_pos, ): - """Faithful reference start_pos==0 path (incl. remainder). Returns (entries[ncomp,d], kv_state, score_state). - - entries have rope applied; kv_state/score_state carry the partial window for the decode path. - """ - overlap = ratio == 4 - coff = 2 if overlap else 1 - d = head_dim - s = x.shape[0] - dtype = x.dtype - xf = x.float() - kv = F.linear(xf, wkv_w.float()) # [s, coff*d] - score = F.linear(xf, wgate_w.float()) # [s, coff*d] - ape = ape.float() - kv_state, score_state = _build_state_from_kv_score(kv, score, ape, ratio, head_dim) - sglang_state_pool = state_pool - try: - comp, sglang_state_pool = _sglang_prefill_from_kv_score( - kv, - score, - norm_w, - ape, - ratio, - head_dim, - cos_table, - sin_table, - eps, - dtype, - state_pool=sglang_state_pool, - ) - if comp is not None: - if return_state_pool: - return comp, kv_state, score_state, sglang_state_pool - return comp, kv_state, score_state - except Exception as exc: - _warn_sglang_fallback(exc) - - should_compress = s >= ratio - remainder = s % ratio - cutoff = s - remainder - if remainder > 0: - kv = kv[:cutoff] - score = score[:cutoff] - if not should_compress: - comp = x.new_zeros(0, head_dim) - if return_state_pool: - return comp, kv_state, score_state, sglang_state_pool - return comp, kv_state, score_state - nwin = cutoff // ratio - kvw = kv.view(nwin, ratio, coff * d) - scw = score.view(nwin, ratio, coff * d) + ape - if overlap: - kvw = _overlap_transform(kvw, ratio, d, 0.0) - scw = _overlap_transform(scw, ratio, d, float("-inf")) - comp = (kvw * torch.softmax(scw, dim=1)).sum(dim=1) # [nwin, d] fp32 - comp = _rmsnorm(comp.to(dtype), norm_w, eps) - pos = torch.arange(nwin, device=x.device) * ratio - comp_rope = apply_rotary_emb(comp[:, -rope_dim:], cos_table[pos], sin_table[pos]) - comp = torch.cat([comp[:, :-rope_dim], comp_rope], dim=1) - if return_state_pool: - return comp, kv_state, score_state, sglang_state_pool - return comp, kv_state, score_state - - -def compressor_decode_step( - x_new, + mod, _ = _load_sglang_compressor() + kv_score = _project_kv_score(x_new, wkv_w, wgate_w) + pool = _state_pool_view(state_pool) + seq_lens = (start_pos + 1).to(torch.int32).contiguous() + plan = mod.CompressorDecodePlan(ratio, seq_lens) + out = mod.compress_forward( + pool, + kv_score, + _sglang_ape(ape.float(), ratio, head_dim), + b_req_idx.to(torch.int32).contiguous(), + plan, + head_dim=head_dim, + compress_ratio=ratio, + ) + should_compress = (seq_lens % ratio) == 0 + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + return out.to(x_new.dtype), should_compress + + +# ---------------------------------------------------------------------------- paged state (c4) +# 与 sglang srt compressor 的 paged 路径同构(compress_old 内核 + 分组槽 indices + overlap +# extra_data): state 槽位由 swa 槽位算术派生(翻译③ state_loc = page*ring + swa_loc%ring, +# 分组槽 = state_loc//ratio),state 随 swa 页生灭,radix 命中零拷贝续算。 + + +def paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: + """state 池行数 = 页数*ring + ring(HOLD 页) + 1(哨兵行),向上取整到 ratio 整除 + (分组视图 [-1, ratio, last_dim] 需要)。与 sglang CompressStatePool 的 _size 公式一致。""" + rows = num_swa_pages * ring + ring + 1 + return (rows + ratio - 1) // ratio * ratio + + +def init_paged_state_pool(buffer: torch.Tensor) -> None: + """末行为哨兵: kv 半边置 0、score 半边置 -inf(KVAndScore.clear 语义)。其余行无需初始化 + (内核在组起点覆写)。buffer: [rows, 2*coff*head_dim] fp32。""" + half = buffer.shape[-1] // 2 + buffer[-1, :half].zero_() + buffer[-1, half:].fill_(float("-inf")) + return + + +def _paged_state_group_slot(req_to_token, full_to_swa, b_req_idx, positions, page_size, ring, ratio): + """位置 -> state 分组槽(= sglang create_paged_compressor_data.get_raw_loc): + state_loc = (swa_loc//page)*ring + swa_loc%ring; 分组槽 = state_loc//ratio。 + 负位置按 sglang 语义 mask 到 0;已出窗(swa_loc<0)的位置落到 -1(哨兵行,score=-inf)。""" + positions = positions.masked_fill(positions < 0, 0) + full = req_to_token[b_req_idx.long(), positions] + swa_loc = full_to_swa[full.long()].long() + state_loc = torch.div(swa_loc, page_size, rounding_mode="floor") * ring + swa_loc % ring + state_loc = torch.where(swa_loc < 0, torch.full_like(state_loc, -1), state_loc) + return torch.div(state_loc, ratio, rounding_mode="floor").to(torch.int32) + + +def paged_decode_state_slots( + req_to_token, + full_to_swa, + b_req_idx, + b_seq_len, + page_size: int, + ring: int, + ratio: int, + hold_req_id: int, + num_swa_pages: int, +): + """decode 步的 state 分组槽(写槽 = 当前组 clip_down(seq-1) 的槽,overlap 伙伴 = 前一组)。 + 纯张量算术(prep 已写本步 req_to_token),图安全。padding(HOLD)行重定向到 HOLD 页的 + state 槽,隔离其垃圾累加。""" + seq = b_seq_len.long() + write_positions = torch.div(seq - 1, ratio, rounding_mode="floor") * ratio + write_slot = _paged_state_group_slot(req_to_token, full_to_swa, b_req_idx, write_positions, page_size, ring, ratio) + overlap_slot = _paged_state_group_slot( + req_to_token, full_to_swa, b_req_idx, write_positions - ratio, page_size, ring, ratio + ) + hold_slot = num_swa_pages * ring // ratio # HOLD 页区域([pages*ring, pages*ring+ring))的首个分组槽 + is_hold = b_req_idx.long() == hold_req_id + write_slot = torch.where(is_hold, torch.full_like(write_slot, hold_slot), write_slot) + overlap_slot = torch.where(is_hold, torch.full_like(overlap_slot, hold_slot), overlap_slot) + return write_slot, overlap_slot + + +def paged_prefill_compress_data(req_to_token, full_to_swa, req_idx: int, ready_len: int, seq_len: int, ring: int): + """单请求 prefill chunk 的 (indices, extra_data, plan): 与 sglang 同走 + triton_create_paged_compress_data(按请求产出,内核经 plan 逐 token 步进)。仅 c4(overlap)。 + 三者都与层无关,同一 forward 内可跨全部 c4 层复用。""" + mod, _ = _load_sglang_compressor() + fn = _load_paged_compress_data_fn() + device = req_to_token.device + n_new = seq_len - ready_len + write_loc, extra_data = fn( + compress_ratio=4, + is_overlap=True, + swa_page_size=128, + ring_size=ring, + req_pool_indices=torch.tensor([req_idx], device=device, dtype=torch.int64), + seq_lens=torch.tensor([seq_len], device=device, dtype=torch.int64), + extend_seq_lens=torch.tensor([n_new], device=device, dtype=torch.int64), + req_to_token=req_to_token, + full_to_swa_index_mapping=full_to_swa, + ) + plan = mod.CompressorPrefillPlan.generate( + 4, + n_new, + torch.tensor([seq_len], dtype=torch.int64), + torch.tensor([n_new], dtype=torch.int64), + device, + ) + return write_loc, extra_data, plan + + +def compressor_paged_prefill( + x, wkv_w, wgate_w, norm_w, ape, - ratio, head_dim, - rope_dim, cos_table, sin_table, eps, - kv_state, - score_state, - start_pos, - state_pool=None, + state_buffer, + compress_data, + ready_len, + seq_len, ): - """Faithful reference start_pos>0 path for one new token. Mutates kv_state/score_state in place. - Returns the new compressed entry [d] (rope applied) when a window completes, else None. - """ - overlap = ratio == 4 - d = head_dim - dtype = x_new.dtype - try: - entry, handled = _sglang_decode_step_from_state_pool( - x_new, - wkv_w, - wgate_w, - norm_w, - ape, - ratio, - head_dim, - cos_table, - sin_table, - eps, - start_pos, - state_pool, - ) - if handled: - return entry - except Exception as exc: - _warn_sglang_fallback(exc) - - xf = x_new.float().view(-1) # [dim] - kv = F.linear(xf, wkv_w.float()) # [coff*d] - score = F.linear(xf, wgate_w.float()) + ape.float()[start_pos % ratio] # [coff*d] - should_compress = (start_pos + 1) % ratio == 0 - if overlap: - kv_state[ratio + start_pos % ratio] = kv - score_state[ratio + start_pos % ratio] = score - if should_compress: - kv_cat = torch.cat([kv_state[:ratio, :d], kv_state[ratio:, d:]], dim=0) # [2*ratio, d] - sc_cat = torch.cat([score_state[:ratio, :d], score_state[ratio:, d:]], dim=0) - entry = (kv_cat * torch.softmax(sc_cat, dim=0)).sum(dim=0) # [d] - kv_state[:ratio] = kv_state[ratio:] - score_state[:ratio] = score_state[ratio:] - else: - kv_state[start_pos % ratio] = kv - score_state[start_pos % ratio] = score - if should_compress: - entry = (kv_state * torch.softmax(score_state, dim=0)).sum(dim=0) # [d] - if not should_compress: - return None - return _finish_entry( - entry, - norm_w, - ape, - rope_dim, - cos_table, - sin_table, - start_pos + 1 - ratio, - eps, - dtype, + """单请求 prefill/extend chunk(c4 paged): x [n_new, dim] 为位置 [ready, seq) 的 hidden, + state 写到 swa 派生的分组槽(compress_data 来自 paged_prefill_compress_data,跨层复用)。 + 返回本 chunk 完结组的压缩条目 [seq//4 - ready//4, head_dim](rope 已施加)。""" + mod, _ = _load_sglang_compressor() + ratio = 4 + kv_score = _project_kv_score(x, wkv_w, wgate_w) + pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) + write_loc, extra_data, plan = compress_data + out = mod.compress_forward( + pool, + kv_score, + _sglang_ape(ape.float(), ratio, head_dim), + write_loc, + plan, + head_dim=head_dim, + compress_ratio=ratio, + extra_data=extra_data, ) + ncomp = seq_len // ratio - ready_len // ratio + if ncomp == 0: + return x.new_zeros(0, head_dim) + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() + return out.index_select(0, ragged_ids).to(x.dtype) -def compressor_decode_step_batch( +def compressor_paged_decode_batch( x_new, wkv_w, wgate_w, norm_w, ape, - ratio, head_dim, - rope_dim, cos_table, sin_table, eps, - state_all, - b_req_idx, - start_pos, + state_buffer, + write_slot, + overlap_slot, + b_seq_len, ): - """Graph-safe batch decode compressor step. - - Mutates ``state_all`` for the selected request rows and returns one candidate - entry per batch row plus a boolean mask telling which rows closed a - compression window. - """ - overlap = ratio == 4 - d = head_dim - dtype = x_new.dtype - req = b_req_idx.long() - pos = start_pos.long() - pos_mod = pos % ratio - - xf = x_new.float() - kv = F.linear(xf, wkv_w.float()) - score = F.linear(xf, wgate_w.float()) + ape.float().index_select(0, pos_mod) - - kv_state = state_all[req, 0].clone() - score_state = state_all[req, 1].clone() - row = pos_mod + (ratio if overlap else 0) - batch_ids = torch.arange(x_new.shape[0], device=x_new.device) - kv_state[batch_ids, row] = kv - score_state[batch_ids, row] = score - - should_compress = ((pos + 1) % ratio) == 0 - if overlap: - kv_cat = torch.cat([kv_state[:, :ratio, :d], kv_state[:, ratio:, d:]], dim=1) - score_cat = torch.cat([score_state[:, :ratio, :d], score_state[:, ratio:, d:]], dim=1) - entry = (kv_cat * torch.softmax(score_cat, dim=1)).sum(dim=1) - shifted_kv_state = kv_state.clone() - shifted_score_state = score_state.clone() - shifted_kv_state[:, :ratio] = kv_state[:, ratio:] - shifted_score_state[:, :ratio] = score_state[:, ratio:] - kv_state = torch.where(should_compress.view(-1, 1, 1), shifted_kv_state, kv_state) - score_state = torch.where(should_compress.view(-1, 1, 1), shifted_score_state, score_state) - else: - entry = (kv_state * torch.softmax(score_state, dim=1)).sum(dim=1) - - state_all[req, 0] = kv_state - state_all[req, 1] = score_state - - entry = _rmsnorm(entry.to(dtype), norm_w, eps) - comp_pos = torch.clamp(pos + 1 - ratio, min=0) - entry_rope = apply_rotary_emb(entry[:, -rope_dim:], cos_table[comp_pos], sin_table[comp_pos]) - entry = torch.cat([entry[:, :-rope_dim], entry_rope], dim=1) - return entry, should_compress + """批量 decode 一步(c4 paged): state 槽位为 swa 派生分组槽(paged_decode_state_slots, + 可跨层复用)。返回 (entries [bs, head_dim], should_compress [bs])。""" + mod, _ = _load_sglang_compressor() + ratio = 4 + kv_score = _project_kv_score(x_new, wkv_w, wgate_w) + pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) + seq_lens = b_seq_len.to(torch.int32).contiguous() + plan = mod.CompressorDecodePlan(ratio, seq_lens) + out = mod.compress_forward( + pool, + kv_score, + _sglang_ape(ape.float(), ratio, head_dim), + write_slot, + plan, + head_dim=head_dim, + compress_ratio=ratio, + extra_data=overlap_slot.view(-1, 1), + ) + should_compress = (seq_lens % ratio) == 0 + mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) + return out.to(x_new.dtype), should_compress diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index 78cdb3a3f8..b125e9ed06 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -1,50 +1,75 @@ import torch +try: + import vllm.model_executor.layers.mhc # noqa: F401 +except Exception as e: + raise RuntimeError("DeepSeek-V4 requires vLLM mHC custom ops; failed to import vllm MHC kernels") from e -def _ensure_vllm_mhc_ops(): - try: - import vllm.model_executor.layers.mhc # noqa: F401 - except Exception as e: - raise RuntimeError("DeepSeek-V4 requires vLLM mHC custom ops; failed to import vllm MHC kernels") from e +# vllm DeepseekV4DecoderLayer.hc_post_alpha +HC_POST_ALPHA = 2.0 -def hc_pre(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps, sinkhorn_iters): - """streams:[N, hc*dim] -> (collapsed[N,dim], post[N,hc,1], comb[N,hc,hc]).""" - _ensure_vllm_mhc_ops() - post, comb, collapsed = torch.ops.vllm.mhc_pre( - residual=streams.view(-1, hc_mult, dim).contiguous(), + +def hc_pre(residual, hc_fn, hc_scale, hc_base, rms_eps, hc_eps, sinkhorn_iters, norm_weight, norm_eps): + """Standalone hc_pre for the first layer. residual:[T, hc, dim] -> + (x[T,dim], residual, post_mix[T,hc,1], res_mix[T,hc,hc]); the sub-layer RMSNorm is fused via norm_weight.""" + post_mix, res_mix, x = torch.ops.vllm.mhc_pre_tilelang( + residual=residual, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=rms_eps, + hc_pre_eps=hc_eps, + hc_sinkhorn_eps=hc_eps, + hc_post_mult_value=HC_POST_ALPHA, + sinkhorn_repeat=sinkhorn_iters, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + return x, residual, post_mix, res_mix + + +def hc_fused_post_pre( + x, residual, post_mix, res_mix, hc_fn, hc_scale, hc_base, rms_eps, hc_eps, sinkhorn_iters, norm_weight, norm_eps +): + """hc_post of the previous sub-layer fused with hc_pre of the next one (norm fused too). + Returns (x[T,dim], residual[T,hc,dim], post_mix, res_mix).""" + residual, post_mix, res_mix, x = torch.ops.vllm.mhc_fused_post_pre_tilelang( + x=x, + residual=residual, + post_layer_mix=post_mix, + comb_res_mix=res_mix, fn=hc_fn, hc_scale=hc_scale, hc_base=hc_base, - rms_eps=eps, - hc_pre_eps=eps, - hc_sinkhorn_eps=eps, - hc_post_mult_value=2.0, + rms_eps=rms_eps, + hc_pre_eps=hc_eps, + hc_sinkhorn_eps=hc_eps, + hc_post_mult_value=HC_POST_ALPHA, sinkhorn_repeat=sinkhorn_iters, + norm_weight=norm_weight, + norm_eps=norm_eps, ) - return collapsed, post, comb + return x, residual, post_mix, res_mix -def hc_post(x, residual, post, comb, hc_mult, dim): - """x:[N,dim] sub-layer output, residual:[N, hc*dim] -> [N, hc*dim].""" - _ensure_vllm_mhc_ops() - out = torch.ops.vllm.mhc_post(x, residual.view(-1, hc_mult, dim).contiguous(), post, comb) - return out.reshape(-1, hc_mult * dim) +def hc_post(x, residual, post_mix, res_mix): + """Complete the hc_post left pending by the last layer. -> streams [T, hc, dim].""" + return torch.ops.vllm.mhc_post_tilelang(x, residual, post_mix, res_mix) -def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, eps): +def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, rms_eps, hc_eps): """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim].""" - _ensure_vllm_mhc_ops() out = torch.empty(streams.shape[0], dim, device=streams.device, dtype=streams.dtype) - torch.ops.vllm.hc_head_fused_kernel( + torch.ops.vllm.hc_head_fused_kernel_tilelang( streams.view(-1, hc_mult, dim).contiguous(), hc_fn, hc_scale, hc_base, out, dim, - eps, - eps, + rms_eps, + hc_eps, hc_mult, ) return out diff --git a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py index c23d03afb7..bc95c249f7 100644 --- a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py @@ -1,5 +1,5 @@ from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer -from .hyper_connection import hc_head +from .hyper_connection import hc_head, hc_post from ..infer_struct import DeepseekV4InferStateInfo @@ -8,6 +8,11 @@ class DeepseekV4PostLayerInfer(LlamaPostLayerInfer): def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): cfg = layer_weight.network_config_ + if isinstance(input_embdings, tuple): + # truncated-layer runs (autotune warmup) end before the last layer's _hc_ffn_out + # collapse; finish the pending hc_post here. + streams = hc_post(*input_embdings) + input_embdings = streams.reshape(streams.shape[0], -1) collapsed = hc_head( input_embdings, layer_weight.hc_head_fn_.weight, @@ -15,6 +20,7 @@ def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, l layer_weight.hc_head_base_.weight, cfg["hc_mult"], cfg["hidden_size"], + cfg["rms_norm_eps"], cfg.get("hc_eps", 1e-6), ) return super().token_forward(collapsed, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py index d83e3082b8..b95f5a14a8 100644 --- a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py @@ -8,16 +8,17 @@ class DeepseekV4PreLayerInfer(LlamaPreLayerInfer): """Token embedding, then expand to the hc_mult parallel residual streams [T, hc_mult*hidden].""" - def _embed_and_expand(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): - emb = layer_weight.wte_weight_(input_ids=input_ids, alloc_func=self.alloc_tensor) # [T, hidden] - if self.tp_world_size_ > 1: - all_reduce(emb, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - hc_mult = layer_weight.network_config_["hc_mult"] - t, hidden = emb.shape - return emb.unsqueeze(1).expand(t, hc_mult, hidden).reshape(t, hc_mult * hidden).contiguous() + def __init__(self, network_config): + super().__init__(network_config) + self.hc_mult = network_config["hc_mult"] + return def context_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): - return self._embed_and_expand(input_ids, infer_state, layer_weight) + input_embdings = super().context_forward(input_ids, infer_state, layer_weight) + t, hidden = input_embdings.shape + return input_embdings.unsqueeze(1).expand(t, self.hc_mult, hidden).reshape(t, self.hc_mult * hidden) def token_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): - return self._embed_and_expand(input_ids, infer_state, layer_weight) + input_embdings = super().token_forward(input_ids, infer_state, layer_weight) + t, hidden = input_embdings.shape + return input_embdings.unsqueeze(1).expand(t, self.hc_mult, hidden).reshape(t, self.hc_mult * hidden) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 11209d39fd..6f7c0de3fc 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -2,28 +2,25 @@ import torch.nn.functional as F import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl +from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.distributed.communication_op import all_reduce from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -from .hyper_connection import hc_pre, hc_post +from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post +from .compressor import ( + compressor_prefill_state, + compressor_decode_step_single, + compressor_decode_step_batch, + compressor_paged_prefill, + compressor_paged_decode_batch, + paged_prefill_compress_data, + paged_decode_state_slots, +) from ..triton_kernel.rotary_emb import apply_rotary_emb from ..infer_struct import DeepseekV4InferStateInfo -from .compressor import compressor_prefill_state, compressor_decode_step, compressor_decode_step_batch -from .attention import vllm_sparse_attn_flat class DeepseekV4TransformerLayerInfer(TransformerLayerInferTpl): - """One V4 decoder layer: HC(attn) then HC(ffn). - - The residual is carried as ``hc_mult`` streams flattened to [T, hc_mult*hidden]; each sub-layer - collapses (hc_pre), computes, and re-expands (hc_post). Attention is MLA over a sliding window + - compressed KV with a per-head sink (vLLM FlashMLA sparse); the MoE reuses lightllm's deepgemm FP8 - grouped GEMM driven by V4's custom router (sqrtsoftplus + hash/topk + bias-for-selection). - - Per-request decode state (window KV history + compressed KV + compressor running state) is kept in - DeepseekV4ReqManager so request alloc/free owns its lifetime. - """ - def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) cfg = network_config @@ -43,6 +40,14 @@ def __init__(self, layer_num, network_config): self.window = cfg["sliding_window"] self.compress_ratio = cfg["compress_ratios"][layer_num] self.is_hash = layer_num < cfg["num_hash_layers"] + self.is_last_layer = layer_num == cfg["n_layer"] - 1 + # complex64 rope table for this layer's variant (sliding / compressed); set by + # DeepseekV4TpPartModel._init_to_get_rotary once the tables are built. The full compress + # cos/sin tables (compressor entry rope uses entry positions, not token positions) are + # wired there too. + self.freqs_cis = None + self.cos_compress_table = None + self.sin_compress_table = None self.topk = cfg["num_experts_per_tok"] self.route_scale = cfg["routed_scaling_factor"] self.swiglu_limit = cfg["swiglu_limit"] @@ -60,47 +65,78 @@ def __init__(self, layer_num, network_config): self.indexer_score_scale = self.index_head_dim ** -0.5 self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 - # ------------------------------------------------------------------ forward (HC-wrapped) - def _hc_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw, attn_forward): - residual = streams - collapsed, post, comb = hc_pre( - streams, - lw.hc_attn_fn_.weight, - lw.hc_attn_scale_.weight, - lw.hc_attn_base_.weight, - self.hc_mult, - self.hidden, + # ------------------------------------------------------------------ forward (HC-threaded) + def _hc_attn_in(self, input_embdings, layer_weight): + """Layer input -> attention input (attn_norm fused). First layer gets the raw streams + and runs a standalone hc_pre; later layers get (x, residual, post_mix, res_mix) and fuse + the previous layer's ffn hc_post with this layer's attn hc_pre.""" + if torch.is_tensor(input_embdings): + residual = input_embdings.view(-1, self.hc_mult, self.hidden) + return hc_pre( + residual, + layer_weight.hc_attn_fn_.weight, + layer_weight.hc_attn_scale_.weight, + layer_weight.hc_attn_base_.weight, + self.eps_, + self.hc_eps, + self.sinkhorn_iters, + layer_weight.attn_norm_.weight, + self.eps_, + ) + x, residual, post_mix, res_mix = input_embdings + return hc_fused_post_pre( + x, + residual, + post_mix, + res_mix, + layer_weight.hc_attn_fn_.weight, + layer_weight.hc_attn_scale_.weight, + layer_weight.hc_attn_base_.weight, + self.eps_, self.hc_eps, self.sinkhorn_iters, + layer_weight.attn_norm_.weight, + self.eps_, ) - o = attn_forward(self._att_norm(collapsed, infer_state, lw), infer_state, lw) - streams = hc_post(o, residual, post, comb, self.hc_mult, self.hidden) - - residual = streams - collapsed, post, comb = hc_pre( - streams, - lw.hc_ffn_fn_.weight, - lw.hc_ffn_scale_.weight, - lw.hc_ffn_base_.weight, - self.hc_mult, - self.hidden, + + def _hc_ffn_in(self, x, residual, post_mix, res_mix, layer_weight): + """Attention output -> ffn input (ffn_norm fused): fused attn hc_post + ffn hc_pre.""" + return hc_fused_post_pre( + x, + residual, + post_mix, + res_mix, + layer_weight.hc_ffn_fn_.weight, + layer_weight.hc_ffn_scale_.weight, + layer_weight.hc_ffn_base_.weight, + self.eps_, self.hc_eps, self.sinkhorn_iters, + layer_weight.ffn_norm_.weight, + self.eps_, ) - f = self._ffn(self._ffn_norm(collapsed, infer_state, lw), infer_state, lw) - return hc_post(f, residual, post, comb, self.hc_mult, self.hidden) - - def context_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw): - return self._hc_forward(streams, infer_state, lw, self.context_attention_forward) - def token_forward(self, streams, infer_state: DeepseekV4InferStateInfo, lw): - return self._hc_forward(streams, infer_state, lw, self.token_attention_forward) - - def _att_norm(self, x, infer_state: DeepseekV4InferStateInfo, lw): - return lw.attn_norm_(x, eps=self.eps_) - - def _ffn_norm(self, x, infer_state: DeepseekV4InferStateInfo, lw): - return lw.ffn_norm_(x, eps=self.eps_) + def _hc_ffn_out(self, x, residual, post_mix, res_mix): + """Mid layers leave the ffn hc_post pending for the next layer's fused post+pre; the last + layer completes it and hands the flat streams [T, hc_mult*hidden] back to the model loop.""" + if not self.is_last_layer: + return x, residual, post_mix, res_mix + streams = hc_post(x, residual, post_mix, res_mix) + return streams.reshape(streams.shape[0], -1) + + def context_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): + x, residual, post_mix, res_mix = self._hc_attn_in(input_embdings, layer_weight) + x = self.context_attention_forward(x, infer_state, layer_weight) + x, residual, post_mix, res_mix = self._hc_ffn_in(x, residual, post_mix, res_mix, layer_weight) + x = self._ffn(x, infer_state, layer_weight) + return self._hc_ffn_out(x, residual, post_mix, res_mix) + + def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): + x, residual, post_mix, res_mix = self._hc_attn_in(input_embdings, layer_weight) + x = self.token_attention_forward(x, infer_state, layer_weight) + x, residual, post_mix, res_mix = self._hc_ffn_in(x, residual, post_mix, res_mix, layer_weight) + x = self._ffn(x, infer_state, layer_weight) + return self._hc_ffn_out(x, residual, post_mix, res_mix) # ------------------------------------------------------------------ shared projections / cache def _select_rope(self, infer_state: DeepseekV4InferStateInfo): @@ -108,20 +144,18 @@ def _select_rope(self, infer_state: DeepseekV4InferStateInfo): return infer_state.position_cos_compress, infer_state.position_sin_compress return infer_state.position_cos_sliding, infer_state.position_sin_sliding - def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, lw): + def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + from sglang.jit_kernel.dsv4 import fused_q_norm_rope + cos_tok, sin_tok = self._select_rope(infer_state) T = x.shape[0] - qa = lw.q_norm_(lw.wq_a_.mm(x), eps=self.eps_) - q = lw.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim).float() - q = (q * torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps_)).to(x.dtype) - q = torch.cat( - [ - q[..., : -self.rope_dim], - apply_rotary_emb(q[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1)), - ], - dim=-1, - ) - kv = lw.kv_norm_(lw.wkv_.mm(x), eps=self.eps_) + qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(x), eps=self.eps_) + q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim) + # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, + # fused in one sglang dsv4 jit kernel (fp32 norm/rotation, bf16 in between -- same as eager). + q = torch.empty_like(q_in) + fused_q_norm_rope(q_in, q, self.eps_, self.freqs_cis, infer_state.position_ids) + kv = layer_weight.kv_norm_(layer_weight.wkv_.mm(x), eps=self.eps_) kv = torch.cat( [ kv[:, : -self.rope_dim], @@ -131,12 +165,12 @@ def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, lw): ) return q, kv, qa, cos_tok, sin_tok - def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, lw): + def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight): # o: [T, tp_q_heads, head_dim] after inverse rope -> grouped low-rank O -> [T, hidden] T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] - o = lw.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] - o = lw.wo_b_.mm(o) + o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] + o = layer_weight.wo_b_.mm(o) if self.tp_world_size_ > 1: all_reduce(o, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) return o @@ -155,813 +189,400 @@ def _inv_rope(self, o, cos_tok, sin_tok): dim=-1, ) - def _post_cache_kv( - self, cache_kv, infer_state: DeepseekV4InferStateInfo, lw, req_idx=None, start_pos=None, mem_index=None - ): - if req_idx is None or start_pos is None or mem_index is None: - raise RuntimeError("DeepSeek-V4 cache write requires req_idx, start_pos, and mem_index") - positions = torch.arange( - start_pos, - start_pos + cache_kv.shape[0], - device=mem_index.device, - dtype=torch.long, - ) - infer_state.mem_manager.pack_mla_kv_to_cache( - layer_index=self.layer_num_, - mem_index=mem_index, - kv=cache_kv.reshape(cache_kv.shape[0], 1, cache_kv.shape[-1]), - req_idx=req_idx, - positions=positions, + # ------------------------------------------------------------------ compressor / indexer + def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + if self.compress_ratio != 4: + return None, None + cos_tok = infer_state.position_cos_compress + sin_tok = infer_state.position_sin_compress + idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_heads, self.index_head_dim) + idx_q = torch.cat( + [ + idx_q[..., : -self.rope_dim], + apply_rotary_emb(idx_q[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1)), + ], + dim=-1, ) - return + idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale + return idx_q, idx_weight - def _get_compressor_state(self, infer_state: DeepseekV4InferStateInfo, req): - cstate_kv, cstate_score = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) - state = { - "cstate_kv": cstate_kv, - "cstate_score": cstate_score, - } - if self.compress_ratio == 4: - idx_state = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) - state["idx_cstate_kv"] = idx_state[req, 0] - state["idx_cstate_score"] = idx_state[req, 1] - return state + def _gather_compress_slots(self, infer_state: DeepseekV4InferStateInfo, req, entry_start, entry_count): + """组末 token 的 full 槽位 -> 压缩槽(条目 [entry_start, entry_start+entry_count))。 + 槽位已由 prep 阶段(prepare_*_compress_slots)分配并 scatter 进 full_to_c4/c128_indexs。""" + ratio = self.compress_ratio + mem = infer_state.mem_manager + mapping = mem.full_to_c4_indexs if ratio == 4 else mem.full_to_c128_indexs + last = entry_start + entry_count + ends = infer_state.req_manager.req_to_token_indexs[req, ratio - 1 : last * ratio : ratio][entry_start:] + return mapping[ends.long()] def _write_compressed_kv(self, infer_state: DeepseekV4InferStateInfo, req, entry_start, comp): - slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, comp.shape[0]) - if comp.shape[0] == 0: - return slots - infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) + slots = self._gather_compress_slots(infer_state, req, entry_start, comp.shape[0]) + if comp.shape[0]: + infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) return slots - def _write_c4_indexer_k(self, infer_state: DeepseekV4InferStateInfo, slots, idx_comp): - if idx_comp is None or idx_comp.shape[0] == 0: - return - infer_state.mem_manager.pack_c4_indexer_k_to_cache(self.layer_num_, slots, idx_comp) - return + def _compressor_weights(self, layer_weight, for_indexer: bool): + if for_indexer: + return ( + layer_weight.idx_cmp_wkv_.mm_param.weight, + layer_weight.idx_cmp_wgate_.mm_param.weight, + layer_weight.idx_cmp_norm_.weight, + layer_weight.idx_cmp_ape_.weight, + self.index_head_dim, + ) + return ( + layer_weight.compressor_wkv_.mm_param.weight, + layer_weight.compressor_wgate_.mm_param.weight, + layer_weight.compressor_norm_.weight, + layer_weight.compressor_ape_.weight, + self.head_dim, + ) - def _dense_kv_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, start_pos, end_pos): - if end_pos <= start_pos: - return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") - slots = infer_state.req_manager.req_to_token_indexs[req, start_pos:end_pos].long() - return infer_state.mem_manager.gather_mla_kv(self.layer_num_, slots) + def _run_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + """Per-request compressor for the prefill chunk. Runs as part of the deferred attention + func, before the attention metadata gathers the slot mappings. - def _compressed_kv_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, ncomp): - if ncomp == 0: - return torch.empty((0, self.head_dim), dtype=infer_state.mem_manager.dtype, device="cuda") + c4: paged state (swa-page-derived group slots, translation #3) — one fused extend-aware + call per request; the (write_loc, extra_data, plan) tuple is layer-independent and cached + on infer_state across all c4 layers. c128: req-keyed state (zero at every 128 boundary by + construction, nothing cache-resident), original jit paths.""" + if not self.compress_ratio: + return if self.compress_ratio == 4: - slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() + self._run_c4_compressor_prefill(x, infer_state, layer_weight) else: - slots = infer_state.req_manager.req_to_c128_indexs[req, :ncomp].long() - return infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots) - - def _c4_indexer_k_from_cache(self, infer_state: DeepseekV4InferStateInfo, req, ncomp): - if self.compress_ratio != 4 or ncomp == 0: - return None - slots = infer_state.req_manager.req_to_c4_indexs[req, :ncomp].long() - return infer_state.mem_manager.gather_c4_indexer_k(self.layer_num_, slots) - - def _run_sparse_attention_batch(self, q_chunks, kv_chunks, index_chunks, sink): - q_flat = torch.cat(q_chunks, dim=0) - kv_flat = torch.cat(kv_chunks, dim=0) - max_topk = max(t.shape[-1] for t in index_chunks) - topk = torch.full( - (q_flat.shape[0], max_topk), - -1, - dtype=torch.int32, - device=q_flat.device, - ) - offset = 0 - for idx in index_chunks: - rows = idx.shape[0] - topk[offset : offset + rows, : idx.shape[1]] = idx.to(torch.int32) - offset += rows - return vllm_sparse_attn_flat(q_flat, kv_flat, sink, topk, self.softmax_scale) - - # ------------------------------------------------------------------ attention (prefill) - def context_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, lw): - q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) - o = self._context_attention_wrapper_run(q, cache_kv, q_lora, x, infer_state, lw) - return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) - - def _context_attention_wrapper_run(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): - if torch.cuda.is_current_stream_capturing(): - q = q.contiguous() - cache_kv = cache_kv.contiguous() - q_lora = q_lora.contiguous() - x = x.contiguous() - _q = tensor_to_no_ref_tensor(q) - _cache_kv = tensor_to_no_ref_tensor(cache_kv) - _q_lora = tensor_to_no_ref_tensor(q_lora) - _x = tensor_to_no_ref_tensor(x) - - pre_capture_graph = infer_state.prefill_cuda_graph_get_current_capture_graph() - pre_capture_graph.__exit__(None, None, None) - - infer_state.prefill_cuda_graph_create_graph_obj() - infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() - o = torch.empty((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) - _o = tensor_to_no_ref_tensor(o) - - def att_func(new_infer_state: DeepseekV4InferStateInfo): - tmp_o = self._context_attention_kernel(_q, _cache_kv, _q_lora, _x, new_infer_state, lw) - assert tmp_o.shape == _o.shape - _o.copy_(tmp_o) - return - - infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) - return o - - return self._context_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) + self._run_c128_compressor_prefill(x, infer_state, layer_weight) + return - def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): - T = x.shape[0] - sink = lw.attn_sink_.weight - o = x.new_empty(T, self.tp_q_heads, self.head_dim) + def _run_c4_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + rm = infer_state.req_manager + mem = infer_state.mem_manager + wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) + iwkv, iwgate, inorm, iape, _ = self._compressor_weights(layer_weight, for_indexer=True) + state_buf = mem.get_c4_state_buffer(self.layer_num_) + idx_state_buf = mem.get_c4_indexer_state_buffer(self.layer_num_) + data_cache = getattr(infer_state, "_dsv4_c4_prefill_data", None) + if data_cache is None: + data_cache = {} + infer_state._dsv4_c4_prefill_data = data_cache b_req = infer_state.b_req_idx.tolist() starts = infer_state.b_q_start_loc.tolist() lens = infer_state.b_q_seq_len.tolist() ready_lens = infer_state.b_ready_cache_len.tolist() - idx_q, idx_weight = self._indexer_q_weight( - x, - q_lora, - infer_state.position_cos_compress, - infer_state.position_sin_compress, - lw, - ) - q_chunks = [] - kv_chunks = [] - index_chunks = [] - out_ranges = [] - kv_offset = 0 - hold_req = infer_state.req_manager.HOLD_REQUEST_ID for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): - if req == hold_req: - o[st : st + ln].zero_() + if req == rm.HOLD_REQUEST_ID or ln == 0: continue - q_r = q[st : st + ln] - cache_kv_r = cache_kv[st : st + ln] + seq_len = ready_len + ln + data = data_cache.get(req) + if data is None: + data = paged_prefill_compress_data( + rm.req_to_token_indexs, mem.full_to_swa_indexs, req, ready_len, seq_len, ring=8 + ) + data_cache[req] = data x_r = x[st : st + ln] - idx_q_r = None if idx_q is None else idx_q[st : st + ln] - idx_weight_r = None if idx_weight is None else idx_weight[st : st + ln] - kv_all, dense_base, n_window, ncomp, idx_comp = self._gather_prefill( - x_r, cache_kv_r, req, ready_len, lw, infer_state - ) - ti = self._topk_idxs_prefill( - ln, - dense_base, - n_window, - ncomp, - x.device, - ready_len, - idx_q_r, - idx_comp, - idx_weight_r, - infer_state, - )[0] - ti = torch.where(ti >= 0, ti + kv_offset, ti).to(torch.int32) - q_chunks.append(q_r) - kv_chunks.append(kv_all) - index_chunks.append(ti) - out_ranges.append((st, ln)) - kv_offset += kv_all.shape[0] - self._post_cache_kv( - cache_kv_r, - infer_state, - lw, - req_idx=req, - start_pos=ready_len, - mem_index=infer_state.mem_index[st : st + ln], - ) - if q_chunks: - attn_out = self._run_sparse_attention_batch(q_chunks, kv_chunks, index_chunks, sink) - out_offset = 0 - for st, ln in out_ranges: - o[st : st + ln] = attn_out[out_offset : out_offset + ln] - out_offset += ln - return o - - def _gather_prefill(self, x_r, kv_r, req, ready_len, lw, infer_state: DeepseekV4InferStateInfo): - ln = kv_r.shape[0] - idx_comp = None - if ready_len > 0: - return self._gather_prefill_extend(x_r, kv_r, req, ready_len, lw, infer_state) - if self.compress_ratio: - cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) - comp, ks, ss, cstate_pool = compressor_prefill_state( + comp = compressor_paged_prefill( x_r, - lw.compressor_wkv_.mm_param.weight, - lw.compressor_wgate_.mm_param.weight, - lw.compressor_norm_.weight, - lw.compressor_ape_.weight, - self.compress_ratio, + wkv, + wgate, + norm, + ape, self.head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, + self.cos_compress_table, + self.sin_compress_table, self.eps_, - return_state_pool=True, - state_pool=cstate_pool, + state_buf, + data, + ready_len, + seq_len, ) - comp_slots = self._write_compressed_kv(infer_state, req, 0, comp) - cstate_kv, cstate_score = infer_state.req_manager.get_compress_state_for_req(self.layer_num_, req) - cstate_kv.copy_(ks) - cstate_score.copy_(ss) - if self.compress_ratio == 4: - idx_cstate_pool = infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) - idx_comp, idx_ks, idx_ss, idx_cstate_pool = compressor_prefill_state( - x_r, - lw.idx_cmp_wkv_.mm_param.weight, - lw.idx_cmp_wgate_.mm_param.weight, - lw.idx_cmp_norm_.weight, - lw.idx_cmp_ape_.weight, - 4, - self.index_head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, - self.eps_, - return_state_pool=True, - state_pool=idx_cstate_pool, - ) - self._write_c4_indexer_k(infer_state, comp_slots, idx_comp) - idx_state = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) - idx_cstate_kv = idx_state[req, 0] - idx_cstate_score = idx_state[req, 1] - idx_cstate_kv.copy_(idx_ks) - idx_cstate_score.copy_(idx_ss) - ncomp = comp.shape[0] - comp = self._compressed_kv_from_cache(infer_state, req, ncomp) - idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) - return torch.cat([kv_r, comp], dim=0), 0, ln, ncomp, idx_comp - return kv_r, 0, ln, 0, None - - def _gather_prefill_extend(self, x_r, kv_r, req, ready_len, lw, infer_state: DeepseekV4InferStateInfo): - if self.compress_ratio: - state = self._get_compressor_state(infer_state, req) - cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) - idx_cstate_pool = ( - infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) - if self.compress_ratio == 4 - else None + slots = self._write_compressed_kv(infer_state, req, ready_len // 4, comp) + idx_comp = compressor_paged_prefill( + x_r, + iwkv, + iwgate, + inorm, + iape, + self.index_head_dim, + self.cos_compress_table, + self.sin_compress_table, + self.eps_, + idx_state_buf, + data, + ready_len, + seq_len, ) + if idx_comp.shape[0]: + infer_state.mem_manager.pack_indexer_k_to_cache(self.layer_num_, slots, idx_comp) + return - for j in range(x_r.shape[0]): - start_pos = ready_len + j - entry = compressor_decode_step( - x_r[j], - lw.compressor_wkv_.mm_param.weight, - lw.compressor_wgate_.mm_param.weight, - lw.compressor_norm_.weight, - lw.compressor_ape_.weight, + def _run_c128_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + rm = infer_state.req_manager + wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) + b_req = infer_state.b_req_idx.tolist() + starts = infer_state.b_q_start_loc.tolist() + lens = infer_state.b_q_seq_len.tolist() + ready_lens = infer_state.b_ready_cache_len.tolist() + for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): + if req == rm.HOLD_REQUEST_ID: + continue + x_r = x[st : st + ln] + state_pool = rm.get_compress_state_pool_for_req(self.layer_num_, req) + if ready_len == 0: + comp = compressor_prefill_state( + x_r, + wkv, + wgate, + norm, + ape, self.compress_ratio, self.head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, + self.cos_compress_table, + self.sin_compress_table, self.eps_, - state["cstate_kv"], - state["cstate_score"], - start_pos, - state_pool=cstate_pool, + state_pool, ) - if entry is not None: - entry_start = (start_pos + 1) // self.compress_ratio - 1 - slots = self._write_compressed_kv(infer_state, req, entry_start, entry.unsqueeze(0)) - if self.compress_ratio == 4: - idx_entry = compressor_decode_step( + self._write_compressed_kv(infer_state, req, 0, comp) + else: + for j in range(ln): + start_pos = ready_len + j + entry = compressor_decode_step_single( x_r[j], - lw.idx_cmp_wkv_.mm_param.weight, - lw.idx_cmp_wgate_.mm_param.weight, - lw.idx_cmp_norm_.weight, - lw.idx_cmp_ape_.weight, - 4, - self.index_head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, + wkv, + wgate, + norm, + ape, + self.compress_ratio, + self.head_dim, + self.cos_compress_table, + self.sin_compress_table, self.eps_, - state["idx_cstate_kv"], - state["idx_cstate_score"], + state_pool, start_pos, - state_pool=idx_cstate_pool, ) - if idx_entry is not None: - if entry is None: - entry_start = (start_pos + 1) // self.compress_ratio - 1 - slots = infer_state.req_manager.ensure_compress_slots(self.layer_num_, req, entry_start, 1) - self._write_c4_indexer_k(infer_state, slots, idx_entry.unsqueeze(0)) - dense_end = ready_len + x_r.shape[0] - ncomp = dense_end // self.compress_ratio - dense_base = max(0, ready_len - self.window + 1) - cached_dense = self._dense_kv_from_cache(infer_state, req, dense_base, ready_len) - dense = torch.cat([cached_dense, kv_r], dim=0) - comp = self._compressed_kv_from_cache(infer_state, req, ncomp) - idx_comp = self._c4_indexer_k_from_cache(infer_state, req, ncomp) - return ( - torch.cat([dense, comp], dim=0), - dense_base, - dense.shape[0], - ncomp, - idx_comp, - ) - dense_base = max(0, ready_len - self.window + 1) - cached_dense = self._dense_kv_from_cache(infer_state, req, dense_base, ready_len) - dense = torch.cat([cached_dense, kv_r], dim=0) - return ( - dense, - dense_base, - dense.shape[0], - 0, - None, - ) + if entry is not None: + entry_start = (start_pos + 1) // self.compress_ratio - 1 + self._write_compressed_kv(infer_state, req, entry_start, entry.unsqueeze(0)) + return - def _topk_idxs_prefill( - self, - seqlen, - dense_base, - n_window, - ncomp, - device, - base_pos, - idx_q, - idx_comp, - idx_weight, - infer_state: DeepseekV4InferStateInfo, - ): - t = torch.arange(seqlen, device=device) - abs_pos = t + base_pos - offsets = torch.arange(self.window, device=device) - win_abs = abs_pos.unsqueeze(1) - (self.window - 1 - offsets).unsqueeze(0) - valid = (win_abs >= dense_base) & (win_abs < dense_base + n_window) - win = torch.where(valid, win_abs - dense_base, torch.full_like(win_abs, -1)) - if ncomp: - if self.compress_ratio == 4 and ncomp > self.index_topk: - comp = self._indexer_topk(idx_q, idx_comp, idx_weight, abs_pos + 1, n_window, infer_state) - else: - c = torch.arange(ncomp, device=device) - comp = torch.where( - c.unsqueeze(0) < ((abs_pos.unsqueeze(1) + 1) // self.compress_ratio), - (c.unsqueeze(0) + n_window).expand(seqlen, ncomp), - torch.full((seqlen, ncomp), -1, device=device, dtype=torch.long), - ) - return torch.cat([win, comp], dim=1).int().unsqueeze(0) - return win.int().unsqueeze(0) - - def _decode_dense_kv_graph(self, infer_state: DeepseekV4InferStateInfo): - req = infer_state.b_req_idx.long() - seq = infer_state.b_seq_len.long() - B = req.shape[0] - device = infer_state.b_seq_len.device - offsets = torch.arange(self.window, device=device, dtype=torch.long) - win_len = torch.minimum(seq, torch.full_like(seq, self.window)) - start = seq - win_len - pos = start.unsqueeze(1) + offsets.unsqueeze(0) - valid = offsets.unsqueeze(0) < win_len.unsqueeze(1) - hold = infer_state.mem_manager.swa_pool.HOLD_TOKEN_MEMINDEX - safe_pos = torch.where(valid, pos, torch.zeros_like(pos)).long() - full_slots = infer_state.req_manager.req_to_token_indexs[req.unsqueeze(1), safe_pos].long() - swa_slots = infer_state.mem_manager.full_to_swa_indexs[full_slots].long() - slot_valid = valid & (swa_slots >= 0) - swa_slots = torch.where(slot_valid, swa_slots, torch.full_like(swa_slots, hold)) - kv = infer_state.mem_manager.gather_mla_kv_from_swa_slots(self.layer_num_, swa_slots.reshape(-1)) - return kv.view(B, self.window, self.head_dim), valid - - def _decode_all_compressed_kv_graph(self, infer_state: DeepseekV4InferStateInfo, ratio): - req = infer_state.b_req_idx.long() - seq = infer_state.b_seq_len.long() - B = req.shape[0] - device = infer_state.b_seq_len.device - max_comp = max(1, infer_state.max_kv_seq_len // ratio) - offsets = torch.arange(max_comp, device=device, dtype=torch.long) - ncomp = torch.div(seq, ratio, rounding_mode="floor") - valid = offsets.unsqueeze(0) < ncomp.unsqueeze(1) - safe_offsets = torch.where(valid, offsets.unsqueeze(0), torch.zeros_like(offsets).unsqueeze(0)) - if ratio == 4: - table = infer_state.req_manager.req_to_c4_indexs - hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX - else: - table = infer_state.req_manager.req_to_c128_indexs - hold = infer_state.mem_manager.c128_pool.HOLD_TOKEN_MEMINDEX - slots = table[req.unsqueeze(1), safe_offsets].long() - slots = torch.where(valid, slots, torch.full_like(slots, hold)) - kv = infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots.reshape(-1)) - kv = kv.view(B, max_comp, self.head_dim) - if ratio != 4: - return kv, None, valid, ncomp - idx_k = infer_state.mem_manager.gather_c4_indexer_k(self.layer_num_, slots.reshape(-1)) - idx_k = idx_k.view(B, max_comp, self.index_head_dim) - return kv, idx_k, valid, ncomp - - def _decode_c4_topk_graph( - self, idx_q, idx_weight, idx_comp, valid_comp, ncomp, infer_state: DeepseekV4InferStateInfo - ): - scores = torch.einsum("bhd,bnd->bhn", idx_q.float(), idx_comp.float()) - scores = F.relu(scores) * self.indexer_score_scale - index_scores = (scores * idx_weight.unsqueeze(-1)).sum(dim=1) - if self.tp_world_size_ > 1: - all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - index_scores = index_scores.masked_fill(~valid_comp, float("-inf")) - top = index_scores.topk(self.index_topk, dim=-1).indices - valid = top < ncomp.unsqueeze(1) - return torch.where(valid, top, torch.zeros_like(top)), valid + def _run_compressor_decode(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + """Batched decode compressor (cuda-graph safe): state update for every request, cache write + masked to the pool HOLD slot unless this token completes a window. Compressed-cache slots + were pre-allocated by prepare_decode_compress_slots in the prep phase. - def _decode_compressed_candidates_graph(self, idx_q, idx_weight, infer_state: DeepseekV4InferStateInfo): - if self.compress_ratio == 4: - _, idx_comp, valid_all, ncomp = self._decode_all_compressed_kv_graph(infer_state, 4) - top, valid = self._decode_c4_topk_graph(idx_q, idx_weight, idx_comp, valid_all, ncomp, infer_state) - req = infer_state.b_req_idx.long() - slots = infer_state.req_manager.req_to_c4_indexs[req.unsqueeze(1), top].long() - hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX - slots = torch.where(valid, slots, torch.full_like(slots, hold)) - comp = infer_state.mem_manager.gather_compressed_kv(self.layer_num_, slots.reshape(-1)) - return comp.view(req.shape[0], self.index_topk, self.head_dim), valid - comp, _, valid, _ = self._decode_all_compressed_kv_graph(infer_state, 128) - return comp, valid - - def _write_decode_compressed_entry_graph(self, x, infer_state: DeepseekV4InferStateInfo, lw, ratio): + c4: paged state — group slots derived from full_to_swa (translation #3) via pure tensor + ops (graph-safe), shared across all c4 layers per step. c128: req-keyed state.""" + if not self.compress_ratio: + return + rm = infer_state.req_manager + mem = infer_state.mem_manager req = infer_state.b_req_idx - start_pos = infer_state.b_seq_len.long() - 1 + ratio = self.compress_ratio + wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) + if ratio == 4: - state_all = infer_state.req_manager.get_c4_compress_state(self.layer_num_) - table = infer_state.req_manager.req_to_c4_indexs - hold = infer_state.mem_manager.c4_pool.HOLD_TOKEN_MEMINDEX + mapping, hold = mem.full_to_c4_indexs, mem.c4_pool.HOLD_TOKEN_MEMINDEX + slot_meta = getattr(infer_state, "_dsv4_c4_decode_slots", None) + if slot_meta is None: + slot_meta = paged_decode_state_slots( + rm.req_to_token_indexs, + mem.full_to_swa_indexs, + req, + infer_state.b_seq_len, + page_size=128, + ring=8, + ratio=4, + hold_req_id=rm.HOLD_REQUEST_ID, + num_swa_pages=mem.swa_num_pages, + ) + infer_state._dsv4_c4_decode_slots = slot_meta + write_slot, overlap_slot = slot_meta + entry, should = compressor_paged_decode_batch( + x, + wkv, + wgate, + norm, + ape, + self.head_dim, + self.cos_compress_table, + self.sin_compress_table, + self.eps_, + mem.get_c4_state_buffer(self.layer_num_), + write_slot, + overlap_slot, + infer_state.b_seq_len, + ) else: - state_all = infer_state.req_manager.get_c128_compress_state(self.layer_num_) - table = infer_state.req_manager.req_to_c128_indexs - hold = infer_state.mem_manager.c128_pool.HOLD_TOKEN_MEMINDEX + mapping, hold = mem.full_to_c128_indexs, mem.c128_pool.HOLD_TOKEN_MEMINDEX + entry, should = compressor_decode_step_batch( + x, + wkv, + wgate, + norm, + ape, + ratio, + self.head_dim, + self.rope_dim, + self.cos_compress_table, + self.sin_compress_table, + self.eps_, + rm.get_compress_state_pool(self.layer_num_), + req, + infer_state.b_seq_len.long() - 1, + ) - entry, should = compressor_decode_step_batch( - x, - lw.compressor_wkv_.mm_param.weight, - lw.compressor_wgate_.mm_param.weight, - lw.compressor_norm_.weight, - lw.compressor_ape_.weight, - ratio, - self.head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, - self.eps_, - state_all, - req, - start_pos, - ) - entry_idx = torch.clamp(torch.div(infer_state.b_seq_len.long(), ratio, rounding_mode="floor") - 1, min=0) - slots = table[req.long(), entry_idx].long() + should = should & (req != rm.HOLD_REQUEST_ID) + # 本步 token 即组末 token(should 为真时),其 full 槽 = mem_index,映射在 prep 已 scatter。 + slots = mapping[infer_state.mem_index.long()].long() slots = torch.where(should, slots, torch.full_like(slots, hold)) - infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, entry) + mem.pack_compressed_kv_to_cache(self.layer_num_, slots, entry) if ratio == 4: - idx_state_all = infer_state.req_manager.get_c4_indexer_compress_state(self.layer_num_) - idx_entry, idx_should = compressor_decode_step_batch( + iwkv, iwgate, inorm, iape, _ = self._compressor_weights(layer_weight, for_indexer=True) + idx_entry, idx_should = compressor_paged_decode_batch( x, - lw.idx_cmp_wkv_.mm_param.weight, - lw.idx_cmp_wgate_.mm_param.weight, - lw.idx_cmp_norm_.weight, - lw.idx_cmp_ape_.weight, - 4, + iwkv, + iwgate, + inorm, + iape, self.index_head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, + self.cos_compress_table, + self.sin_compress_table, self.eps_, - idx_state_all, - req, - start_pos, + mem.get_c4_indexer_state_buffer(self.layer_num_), + write_slot, + overlap_slot, + infer_state.b_seq_len, ) + idx_should = idx_should & (req != rm.HOLD_REQUEST_ID) idx_slots = torch.where(idx_should, slots, torch.full_like(slots, hold)) - infer_state.mem_manager.pack_c4_indexer_k_to_cache(self.layer_num_, idx_slots, idx_entry) + mem.pack_indexer_k_to_cache(self.layer_num_, idx_slots, idx_entry) return - # ------------------------------------------------------------------ attention (decode) - def token_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, lw): - q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, lw) - if infer_state.is_cuda_graph: - o = self._token_attention_kernel_cuda_graph(q, cache_kv, q_lora, x, infer_state, lw) - else: - o = self._token_attention_kernel(q, cache_kv, q_lora, x, infer_state, lw) - return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, lw) - - def _token_attention_kernel_cuda_graph(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): - sink = lw.attn_sink_.weight - infer_state.mem_manager.pack_decode_mla_kv_to_cache( - self.layer_num_, - infer_state.b_req_idx, - infer_state.b_seq_len, - infer_state.mem_index, - cache_kv.reshape(cache_kv.shape[0], 1, cache_kv.shape[-1]), - ) - idx_q, idx_weight = self._indexer_q_weight( - x, - q_lora, - infer_state.position_cos_compress, - infer_state.position_sin_compress, - lw, - ) - if self.compress_ratio: - self._write_decode_compressed_entry_graph(x, infer_state, lw, self.compress_ratio) + # ------------------------------------------------------------------ attention (prefill) + def context_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + # template hook: write the chunk's packed latent into the swa pool before attention + # reads it back via full_to_swa indices (this custom forward bypasses the tpl path). + self._post_cache_kv(cache_kv, infer_state, layer_weight) + o = self._context_attention_wrapper_run(q, cache_kv, q_lora, x, infer_state, layer_weight) + return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) + + def _context_attention_wrapper_run( + self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight + ): + if torch.cuda.is_current_stream_capturing(): + q = q.contiguous() + cache_kv = cache_kv.contiguous() + q_lora = q_lora.contiguous() + x = x.contiguous() + _q = tensor_to_no_ref_tensor(q) + _cache_kv = tensor_to_no_ref_tensor(cache_kv) + _q_lora = tensor_to_no_ref_tensor(q_lora) + _x = tensor_to_no_ref_tensor(x) - dense_kv, dense_valid = self._decode_dense_kv_graph(infer_state) - B = q.shape[0] - device = q.device - if self.compress_ratio: - comp_kv, comp_valid = self._decode_compressed_candidates_graph(idx_q, idx_weight, infer_state) - kv_all = torch.cat([dense_kv, comp_kv], dim=1) - comp_offsets = torch.arange(comp_kv.shape[1], device=device, dtype=torch.int32) - else: - kv_all = dense_kv - comp_valid = None - comp_offsets = None - - total_k = kv_all.shape[1] - base = torch.arange(B, device=device, dtype=torch.int32).unsqueeze(1) * total_k - dense_offsets = torch.arange(self.window, device=device, dtype=torch.int32) - dense_topk = torch.where( - dense_valid, - base + dense_offsets.unsqueeze(0), - torch.full((B, self.window), -1, device=device, dtype=torch.int32), - ) - if self.compress_ratio: - comp_topk = torch.where( - comp_valid, - base + self.window + comp_offsets.unsqueeze(0), - torch.full((B, comp_kv.shape[1]), -1, device=device, dtype=torch.int32), - ) - topk = torch.cat([dense_topk, comp_topk], dim=1) - else: - topk = dense_topk - return vllm_sparse_attn_flat( - q, - kv_all.reshape(-1, self.head_dim), - sink, - topk, - self.softmax_scale, - already_compact=True, - ) + pre_capture_graph = infer_state.prefill_cuda_graph_get_current_capture_graph() + pre_capture_graph.__exit__(None, None, None) - def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, lw): - B = x.shape[0] # one new token per request - idx_q, idx_weight = self._indexer_q_weight( - x, - q_lora, - infer_state.position_cos_compress, - infer_state.position_sin_compress, - lw, - ) - sink = lw.attn_sink_.weight - b_req = infer_state.b_req_idx.tolist() - seqlens = infer_state.b_seq_len.tolist() - o = x.new_empty(B, self.tp_q_heads, self.head_dim) - hold_req = infer_state.req_manager.HOLD_REQUEST_ID - q_chunks = [] - kv_chunks = [] - index_chunks = [] - out_rows = [] - kv_offset = 0 - for i, (req, seq) in enumerate(zip(b_req, seqlens)): - if req == hold_req: - o[i].zero_() - continue - start_pos = seq - 1 - self._post_cache_kv( - cache_kv[i : i + 1], - infer_state, - lw, - req_idx=req, - start_pos=start_pos, - mem_index=infer_state.mem_index[i : i + 1], - ) - if self.compress_ratio: - stt = self._get_compressor_state(infer_state, req) - cstate_pool = infer_state.req_manager.get_compress_state_pool_for_req(self.layer_num_, req) - e = compressor_decode_step( - x[i], - lw.compressor_wkv_.mm_param.weight, - lw.compressor_wgate_.mm_param.weight, - lw.compressor_norm_.weight, - lw.compressor_ape_.weight, - self.compress_ratio, - self.head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, - self.eps_, - stt["cstate_kv"], - stt["cstate_score"], - start_pos, - state_pool=cstate_pool, - ) - entry_slots = None - if e is not None: - entry_start = (start_pos + 1) // self.compress_ratio - 1 - entry_slots = self._write_compressed_kv(infer_state, req, entry_start, e.unsqueeze(0)) - if self.compress_ratio == 4: - idx_cstate_pool = infer_state.req_manager.get_c4_indexer_state_pool_for_req(self.layer_num_, req) - idx_e = compressor_decode_step( - x[i], - lw.idx_cmp_wkv_.mm_param.weight, - lw.idx_cmp_wgate_.mm_param.weight, - lw.idx_cmp_norm_.weight, - lw.idx_cmp_ape_.weight, - 4, - self.index_head_dim, - self.rope_dim, - infer_state.cos_compress_table, - infer_state.sin_compress_table, - self.eps_, - stt["idx_cstate_kv"], - stt["idx_cstate_score"], - start_pos, - state_pool=idx_cstate_pool, - ) - if idx_e is not None: - if entry_slots is None: - entry_start = (start_pos + 1) // self.compress_ratio - 1 - entry_slots = infer_state.req_manager.ensure_compress_slots( - self.layer_num_, req, entry_start, 1 - ) - self._write_c4_indexer_k(infer_state, entry_slots, idx_e.unsqueeze(0)) - win_start = max(0, seq - self.window) - win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) - comp_kv = self._compressed_kv_from_cache(infer_state, req, seq // self.compress_ratio) - idx_comp = self._c4_indexer_k_from_cache(infer_state, req, comp_kv.shape[0]) - kv_all = torch.cat([win_kv, comp_kv], dim=0) - else: - win_start = max(0, seq - self.window) - win_kv = self._dense_kv_from_cache(infer_state, req, win_start, seq) - kv_all = win_kv - comp_kv = None - idx_comp = None - ti = self._topk_idxs_decode( - win_kv.shape[0], - comp_kv, - None if idx_q is None else idx_q[i : i + 1], - idx_comp, - None if idx_weight is None else idx_weight[i : i + 1], - seq, - x.device, - infer_state, - )[0, 0] - ti = torch.where(ti >= 0, ti + kv_offset, ti).view(1, -1).to(torch.int32) - q_chunks.append(q[i : i + 1]) - kv_chunks.append(kv_all) - index_chunks.append(ti) - out_rows.append(i) - kv_offset += kv_all.shape[0] - if q_chunks: - attn_out = self._run_sparse_attention_batch(q_chunks, kv_chunks, index_chunks, sink) - for row, row_out in zip(out_rows, attn_out): - o[row] = row_out - return o + infer_state.prefill_cuda_graph_create_graph_obj() + infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() + o = torch.empty((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) + _o = tensor_to_no_ref_tensor(o) - def _indexer_q_weight(self, x, qa, cos_tok, sin_tok, lw): - if self.compress_ratio != 4: - return None, None - idx_q = lw.idx_wq_b_.mm(qa).view(x.shape[0], self.tp_index_heads, self.index_head_dim) - idx_q = torch.cat( - [ - idx_q[..., : -self.rope_dim], - apply_rotary_emb( - idx_q[..., -self.rope_dim :], - cos_tok.unsqueeze(1), - sin_tok.unsqueeze(1), - ), - ], - dim=-1, - ) - idx_weight = lw.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale - return idx_q, idx_weight + def att_func(new_infer_state: DeepseekV4InferStateInfo): + tmp_o = self._context_attention_kernel(_q, _cache_kv, _q_lora, _x, new_infer_state, layer_weight) + assert tmp_o.shape == _o.shape + _o.copy_(tmp_o) + return - def _indexer_topk( - self, idx_q, idx_comp, idx_weight, positions_1based, offset, infer_state: DeepseekV4InferStateInfo - ): - ncomp = idx_comp.shape[0] - k = min(self.index_topk, ncomp) - if k == 0: - return torch.empty((idx_q.shape[0], 0), device=idx_q.device, dtype=torch.long) - - top_chunks = [] - heads = max(1, idx_q.shape[1]) - max_score_elems = 16 * 1024 * 1024 - chunk_size = max(1, min(idx_q.shape[0], max_score_elems // max(1, heads * ncomp))) - for start in range(0, idx_q.shape[0], chunk_size): - end = min(idx_q.shape[0], start + chunk_size) - scores = torch.einsum("thd,nd->thn", idx_q[start:end].float(), idx_comp.float()) - scores = F.relu(scores) * self.indexer_score_scale - index_scores = (scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1) - if self.tp_world_size_ > 1: - all_reduce( - index_scores, - op=dist.ReduceOp.SUM, - group=infer_state.dist_group, - async_op=False, - ) - causal_threshold = positions_1based[start:end] // 4 - top_chunks.append(self._indexer_topk_kernel(index_scores, causal_threshold, k)) - top = torch.cat(top_chunks, dim=0) - valid = top >= 0 - return torch.where(valid, top + offset, torch.full_like(top, -1)) - - def _indexer_topk_kernel(self, index_scores, causal_threshold, topk): - if index_scores.is_cuda: - try: - import vllm._C # noqa: F401 - - scores = index_scores.contiguous() - lengths = causal_threshold.to(torch.int32).contiguous() - starts = torch.zeros_like(lengths, dtype=torch.int32) - top = torch.empty((scores.shape[0], topk), dtype=torch.int32, device=scores.device) - torch.ops._C.top_k_per_row_prefill( - scores, - starts, - lengths, - top, - scores.shape[0], - scores.stride(0), - scores.stride(1), - topk, - ) - return top.long() - except Exception: - pass + infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) + return o - entry_indices = torch.arange(index_scores.shape[1], device=index_scores.device) - index_scores = index_scores.masked_fill( - entry_indices.unsqueeze(0) >= causal_threshold.unsqueeze(1), float("-inf") + return self._context_attention_kernel(q, cache_kv, q_lora, x, infer_state, layer_weight) + + def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + self._run_compressor_prefill(x, infer_state, layer_weight) + idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + att_control = AttControl( + nsa_prefill=True, + nsa_prefill_dict={ + "flashmla_kvcache": True, + "layer_index": self.layer_num_, + "compress_ratio": self.compress_ratio, + "head_dim_v": self.head_dim, + "softmax_scale": self.softmax_scale, + "cache_kv": cache_kv, + "q_lora": q_lora, + "hidden_states": x, + "attn_sink": layer_weight.attn_sink_.weight, + "idx_q": idx_q, + "idx_weight": idx_weight, + "index_topk": self.index_topk, + "indexer_score_scale": self.indexer_score_scale, + "tp_world_size": self.tp_world_size_, + }, + ) + return infer_state.prefill_att_state.prefill_att( + q=q, + k=infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_), + v=None, + att_control=att_control, ) - top = index_scores.topk(topk, dim=-1).indices - valid = top < causal_threshold.unsqueeze(1) - return torch.where(valid, top, torch.full_like(top, -1)) - - def _topk_idxs_decode( - self, - win_len, - comp_kv, - idx_q, - idx_comp, - idx_weight, - seq_len, - device, - infer_state: DeepseekV4InferStateInfo, - ): - win = torch.arange(win_len, device=device, dtype=torch.long) - if comp_kv is None or comp_kv.shape[0] == 0: - return win.view(1, 1, -1).int() - ncomp = comp_kv.shape[0] - if self.compress_ratio == 4 and ncomp > self.index_topk: - comp = self._indexer_topk( - idx_q, - idx_comp, - idx_weight, - torch.tensor([seq_len], device=device, dtype=torch.long), - win_len, - infer_state, - )[0] - else: - comp = torch.arange(ncomp, device=device, dtype=torch.long) + win_len - return torch.cat([win, comp], dim=0).view(1, 1, -1).int() - # ------------------------------------------------------------------ moe - def _fp4_experts(self, x, weights, indices, lw): - experts = lw.experts_ - if getattr(experts, "moe_backend", None) != "marlin": - err = getattr(experts, "moe_backend_error", "unknown") - raise RuntimeError(f"DeepSeek-V4 FP4 MoE requires vLLM Marlin backend, init_error={err}") - return self._fp4_experts_marlin(x, weights, indices, experts) - - def _fp4_experts_marlin(self, x, weights, indices, experts): - from vllm.model_executor.layers.fused_moe.activation import MoEActivation - from vllm.model_executor.layers.fused_moe.experts.marlin_moe import ( - fused_marlin_moe, + # ------------------------------------------------------------------ attention (decode) + def token_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + self._post_cache_kv(cache_kv, infer_state, layer_weight) + o = self._token_attention_kernel(q, cache_kv, q_lora, x, infer_state, layer_weight) + return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) + + def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + self._run_compressor_decode(x, infer_state, layer_weight) + idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + att_control = AttControl( + nsa_decode=True, + nsa_decode_dict={ + "flashmla_kvcache": True, + "layer_index": self.layer_num_, + "compress_ratio": self.compress_ratio, + "head_dim_v": self.head_dim, + "softmax_scale": self.softmax_scale, + "cache_kv": cache_kv, + "q_lora": q_lora, + "hidden_states": x, + "attn_sink": layer_weight.attn_sink_.weight, + "idx_q": idx_q, + "idx_weight": idx_weight, + "index_topk": self.index_topk, + "indexer_score_scale": self.indexer_score_scale, + "tp_world_size": self.tp_world_size_, + }, ) - from vllm.scalar_type import scalar_types - - return fused_marlin_moe( - hidden_states=x.contiguous(), - w1=experts.marlin_w13, - w2=experts.marlin_w2, - bias1=None, - bias2=None, - w1_scale=experts.marlin_w13_scale, - w2_scale=experts.marlin_w2_scale, - topk_weights=weights.to(torch.float32).contiguous(), - topk_ids=indices.to(torch.long).contiguous(), - quant_type_id=scalar_types.float4_e2m1f.id, - global_num_experts=experts.n_routed_experts, - activation=MoEActivation.SILU, + return infer_state.decode_att_state.decode_att( + q=q, + k=infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_), + v=None, + att_control=att_control, + ) + + # ------------------------------------------------------------------ moe + def _routed_experts(self, x, weights, indices, layer_weight): + return layer_weight.experts_.experts_with_preselected( + input_tensor=x, + topk_weights=weights, + topk_ids=indices, clamp_limit=float(self.swiglu_limit), ) - def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, lw): - gw = lw.gate_weight_.mm_param.weight + def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + gw = layer_weight.gate_weight_.mm_param.weight logits = F.linear(x.float(), gw.float()).contiguous() - weights, indices = self._select_experts(logits, infer_state, lw) - routed = self._fp4_experts(x, weights, indices, lw) - g = lw.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) - u = lw.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) - shared = lw.shared_down_.mm((F.silu(g) * u).to(x.dtype)) - if self.enable_ep_moe and getattr(lw.experts_, "is_ep", False): + weights, indices = self._select_experts(logits, infer_state, layer_weight) + routed = self._routed_experts(x, weights, indices, layer_weight) + g = layer_weight.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) + u = layer_weight.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) + shared = layer_weight.shared_down_.mm((F.silu(g) * u).to(x.dtype)) + if self.enable_ep_moe and getattr(layer_weight.experts_, "is_ep", False): if self.tp_world_size_ > 1: all_reduce( shared, @@ -975,10 +596,10 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, lw): all_reduce(out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) return out - def _select_experts(self, logits, infer_state: DeepseekV4InferStateInfo, lw): - return self._select_experts_vllm(logits, infer_state, lw) + def _select_experts(self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight): + return self._select_experts_vllm(logits, infer_state, layer_weight) - def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, lw): + def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight): from vllm import _custom_ops as ops M = logits.shape[0] @@ -987,13 +608,13 @@ def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, lw hash_indices_table = None indices_dtype = torch.int64 if self.is_hash: - hash_indices_table = lw.gate_tid2eid_.weight + hash_indices_table = layer_weight.gate_tid2eid_.weight if not hash_indices_table.is_contiguous(): hash_indices_table = hash_indices_table.contiguous() indices_dtype = hash_indices_table.dtype input_tokens = infer_state.input_ids.to(dtype=indices_dtype).contiguous() else: - bias = lw.gate_bias_.weight + bias = layer_weight.gate_bias_.weight weights = torch.empty((M, self.topk), dtype=torch.float32, device=logits.device) indices = torch.empty((M, self.topk), dtype=indices_dtype, device=logits.device) diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index cdaaac2cdb..a95299628c 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -1,5 +1,3 @@ -import threading - import torch from lightllm.common.basemodel import TransformerLayerWeight from lightllm.common.basemodel.layer_weights.meta_weights import ( @@ -9,202 +7,16 @@ RMSNormWeight, ParameterWeight, TpAttSinkWeight, + FusedMoeWeight, ) -from lightllm.common.basemodel.layer_weights.meta_weights.base_weight import BaseWeightTpl -from lightllm.common.quantization.registry import QUANTMETHODS -from lightllm.utils.log_utils import init_logger from ..triton_kernel.quant_convert import dequant_fp8_block_to_bf16 -logger = init_logger(__name__) - - -class DeepseekV4FP4ExpertsWeight(BaseWeightTpl): - _marlin_pack_lock = threading.Lock() - - def __init__(self, weight_prefix, n_routed_experts, hidden_size, moe_intermediate_size, data_type): - super().__init__(data_type=data_type) - self.weight_prefix = weight_prefix - self.n_routed_experts = n_routed_experts - self.hidden_size = hidden_size - self.moe_intermediate_size = moe_intermediate_size - self.split_inter_size = moe_intermediate_size // self.tp_world_size_ - self.local_expert_ids = list(range(n_routed_experts)) - self.expert_idx_to_local_idx = {expert_idx: expert_idx for expert_idx in self.local_expert_ids} - self.moe_backend = None - self.moe_backend_error = None - self._marlin_checked = False - self._load_lock = threading.Lock() - self.load_ok = { - name: [False] * n_routed_experts for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale") - } - - def _create_weight(self): - self._ensure_raw_fp4_weight() - - def _ensure_raw_fp4_weight(self): - if hasattr(self, "w1"): - return - device = "cpu" - n = self.n_routed_experts - h = self.hidden_size - inter = self.split_inter_size - self.w1 = torch.empty((n, inter, h // 2), dtype=torch.int8, device=device) - self.w3 = torch.empty((n, inter, h // 2), dtype=torch.int8, device=device) - self.w2 = torch.empty((n, h, inter // 2), dtype=torch.int8, device=device) - self.w1_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) - self.w3_scale = torch.empty((n, inter, h // 32), dtype=torch.float8_e8m0fnu, device=device) - self.w2_scale = torch.empty((n, h, inter // 32), dtype=torch.float8_e8m0fnu, device=device) - - def _copy_expert_weight(self, dst, weight, expert_idx, name, is_down=False): - if is_down: - start = self.tp_rank_ * self.split_inter_size // 2 - end = (self.tp_rank_ + 1) * self.split_inter_size // 2 - src = weight[:, start:end] - else: - start = self.tp_rank_ * self.split_inter_size - end = (self.tp_rank_ + 1) * self.split_inter_size - src = weight[start:end, :] - dst[expert_idx].copy_(src) - self.load_ok[name][expert_idx] = True - - def _copy_expert_scale(self, dst, scale, expert_idx, name, is_down=False): - if is_down: - start = self.tp_rank_ * self.split_inter_size // 32 - end = (self.tp_rank_ + 1) * self.split_inter_size // 32 - src = scale[:, start:end] - else: - start = self.tp_rank_ * self.split_inter_size - end = (self.tp_rank_ + 1) * self.split_inter_size - src = scale[start:end, :] - dst[expert_idx].copy_(src) - self.load_ok[name][expert_idx] = True - - def load_hf_weights(self, weights): - if self._marlin_checked: - return - has_weight = False - for expert_idx in self.local_expert_ids: - prefix = f"{self.weight_prefix}.{expert_idx}" - if ( - f"{prefix}.w1.weight" in weights - or f"{prefix}.w1.scale" in weights - or f"{prefix}.w2.weight" in weights - or f"{prefix}.w2.scale" in weights - or f"{prefix}.w3.weight" in weights - or f"{prefix}.w3.scale" in weights - ): - has_weight = True - break - if not has_weight: - return - - with self._load_lock: - if self._marlin_checked: - return - self._ensure_raw_fp4_weight() - for expert_idx in self.local_expert_ids: - prefix = f"{self.weight_prefix}.{expert_idx}" - w1 = f"{prefix}.w1.weight" - w1_scale = f"{prefix}.w1.scale" - w2 = f"{prefix}.w2.weight" - w2_scale = f"{prefix}.w2.scale" - w3 = f"{prefix}.w3.weight" - w3_scale = f"{prefix}.w3.scale" - if w1 in weights: - self._copy_expert_weight(self.w1, weights[w1], expert_idx, "w1") - if w1_scale in weights: - self._copy_expert_scale(self.w1_scale, weights[w1_scale], expert_idx, "w1_scale") - if w3 in weights: - self._copy_expert_weight(self.w3, weights[w3], expert_idx, "w3") - if w3_scale in weights: - self._copy_expert_scale(self.w3_scale, weights[w3_scale], expert_idx, "w3_scale") - if w2 in weights: - self._copy_expert_weight(self.w2, weights[w2], expert_idx, "w2", is_down=True) - if w2_scale in weights: - self._copy_expert_scale(self.w2_scale, weights[w2_scale], expert_idx, "w2_scale", is_down=True) - if self._raw_load_complete(): - self._try_init_marlin() - - def verify_load(self): - with self._load_lock: - ok = self._raw_load_complete() - if ok and not self._marlin_checked: - self._try_init_marlin() - return ok - - def _raw_load_complete(self): - return all(all(ok_list) for ok_list in self.load_ok.values()) - - def _try_init_marlin(self): - try: - from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( - prepare_moe_mxfp4_layer_for_marlin, - ) - - class _MarlinLayer: - pass - - with self._marlin_pack_lock: - torch.cuda.set_device(self.device_id_) - device = torch.device("cuda", self.device_id_) - layer = _MarlinLayer() - layer.params_dtype = self.data_type_ - w13_cpu, w13_scale_cpu = self._build_w13_weight() - w13 = w13_cpu.to(device=device, non_blocking=True).contiguous() - w2 = self.w2.view(torch.uint8).to(device=device, non_blocking=True).contiguous() - w13_scale = w13_scale_cpu.to(device=device, non_blocking=True).contiguous() - w2_scale = self.w2_scale.to(device=device, non_blocking=True).contiguous() - ( - self.marlin_w13, - self.marlin_w2, - self.marlin_w13_scale, - self.marlin_w2_scale, - _, - _, - ) = prepare_moe_mxfp4_layer_for_marlin(layer, w13, w2, w13_scale, w2_scale, None, None) - del w13_cpu, w13_scale_cpu, w13, w2, w13_scale, w2_scale - self.moe_backend = "marlin" - self._marlin_checked = True - self._release_raw_fp4_weight() - torch.cuda.empty_cache() - logger.info( - "DeepSeek-V4 FP4 experts use vLLM Marlin backend, prefix=%s, rank=%s", - self.weight_prefix, - self.tp_rank_, - ) - except Exception as e: - self.moe_backend_error = repr(e) - raise RuntimeError( - "DeepSeek-V4 FP4 experts require vLLM Marlin backend, " - f"prefix={self.weight_prefix}, rank={self.tp_rank_}, error={self.moe_backend_error}" - ) from e - - def _build_w13_weight(self): - n = self.n_routed_experts - h = self.hidden_size - inter = self.split_inter_size - w13 = torch.empty((n, 2 * inter, h // 2), dtype=torch.uint8, device=self.w1.device) - w13[:, :inter, :].copy_(self.w1.view(torch.uint8)) - w13[:, inter:, :].copy_(self.w3.view(torch.uint8)) - w13_scale = torch.empty((n, 2 * inter, h // 32), dtype=self.w1_scale.dtype, device=self.w1_scale.device) - w13_scale[:, :inter, :].copy_(self.w1_scale) - w13_scale[:, inter:, :].copy_(self.w3_scale) - return w13.contiguous(), w13_scale.contiguous() - - def _release_raw_fp4_weight(self): - for name in ("w1", "w1_scale", "w2", "w2_scale", "w3", "w3_scale"): - if hasattr(self, name): - delattr(self, name) - - class DeepseekV4TransformerLayerWeight(TransformerLayerWeight): """Per-layer weights for DeepSeek-V4-Flash. - The checkpoint stores most linears in FP8 (e4m3 + block-128 ue8m0 scale) and the routed - experts in FP4 (int8-packed e2m1 + group-32 ue8m0 scale). Hopper does not use the SM100 - MegaMoE path here, so routed experts are kept in packed FP4 and temporarily de-quantized only - for selected experts in the correctness-first torch MoE path. + DS4 does not share DS2/DS3.2's ``model.layers.*.self_attn/mlp`` layout. Its attention is + HC + CSA, and routed experts are checkpointed as MXFP4. """ def __init__(self, layer_num, data_type, network_config, quant_cfg=None): @@ -213,7 +25,6 @@ def __init__(self, layer_num, data_type, network_config, quant_cfg=None): def _parse_config(self): cfg = self.network_config_ - self.fp8_quant = QUANTMETHODS.get("deepgemm-fp8w8a8-b128") self.hidden = cfg["hidden_size"] self.n_heads = cfg["num_attention_heads"] self.head_dim = cfg["head_dim"] @@ -238,13 +49,10 @@ def _parse_config(self): assert self.index_n_heads % self.tp_world_size_ == 0 self.prefix = f"layers.{self.layer_num_}" - def _init_weight_names(self): - return - def _init_weight(self): - self._init_attn() + self._init_qkvo() if self.has_compressor: - self._init_compressor(f"{self.prefix}.attn.compressor", self.head_dim, self.compress_ratio) + self._init_compressor() if self.has_indexer: self._init_indexer() self._init_moe() @@ -252,7 +60,7 @@ def _init_weight(self): self._init_hyper_connection() # ------------------------------------------------------------------ attention - def _init_attn(self): + def _init_qkvo(self): p = f"{self.prefix}.attn" # q low-rank (a replicated, b column-parallel over heads), kv single head (replicated) self.wq_a_ = ROWMMWeight( @@ -260,7 +68,7 @@ def _init_attn(self): out_dims=[self.q_lora_rank], weight_names=f"{p}.wq_a.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("wq_a"), tp_rank=0, tp_world_size=1, ) @@ -269,14 +77,14 @@ def _init_attn(self): out_dims=[self.n_heads * self.head_dim], weight_names=f"{p}.wq_b.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("wq_b"), ) self.wkv_ = ROWMMWeight( in_dim=self.hidden, out_dims=[self.head_dim], weight_names=f"{p}.wkv.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("wkv"), tp_rank=0, tp_world_size=1, ) @@ -301,11 +109,15 @@ def _init_attn(self): out_dims=[self.hidden], weight_names=f"{p}.wo_b.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("wo_b"), ) # ------------------------------------------------------------------ compressor / indexer - def _init_compressor(self, prefix, head_dim, ratio): + def _init_compressor(self): + prefix = f"{self.prefix}.attn.compressor" + head_dim = self.head_dim + ratio = self.compress_ratio + coff = 2 if ratio == 4 else 1 # wkv/wgate are bf16 (no scale) and replicated (single KV head). self.compressor_wkv_ = ROWMMWeight( @@ -341,7 +153,7 @@ def _init_indexer(self): out_dims=[self.index_n_heads * self.index_head_dim], weight_names=f"{p}.wq_b.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("idx_wq_b"), ) self.idx_weights_proj_ = ROWMMWeight( in_dim=self.hidden, @@ -406,28 +218,35 @@ def _init_moe(self): out_dims=[self.moe_inter], weight_names=f"{sp}.w1.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("shared_gate"), ) self.shared_up_ = ROWMMWeight( in_dim=self.hidden, out_dims=[self.moe_inter], weight_names=f"{sp}.w3.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("shared_up"), ) self.shared_down_ = COLMMWeight( in_dim=self.moe_inter, out_dims=[self.hidden], weight_names=f"{sp}.w2.weight", data_type=self.data_type_, - quant_method=self.fp8_quant, + quant_method=self.get_quant_method("shared_down"), ) - self.experts_ = DeepseekV4FP4ExpertsWeight( + self.experts_ = FusedMoeWeight( + gate_proj_name="w1", + down_proj_name="w2", + up_proj_name="w3", + e_score_correction_bias_name="", weight_prefix=f"{p}.experts", n_routed_experts=self.n_routed_experts, hidden_size=self.hidden, moe_intermediate_size=self.moe_inter, data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(self.layer_num_, "fused_moe"), + layer_num=self.layer_num_, + network_config=self.network_config_, ) def _init_norm(self): @@ -439,62 +258,68 @@ def _init_norm(self): ) def _init_hyper_connection(self): - for which in ["attn", "ffn"]: - setattr( - self, - f"hc_{which}_fn_", - ParameterWeight( - weight_name=f"{self.prefix}.hc_{which}_fn", - data_type=torch.float32, - weight_shape=(self.mix_hc, self.hc_mult * self.hidden), - ), - ) - setattr( - self, - f"hc_{which}_base_", - ParameterWeight( - weight_name=f"{self.prefix}.hc_{which}_base", data_type=torch.float32, weight_shape=(self.mix_hc,) - ), - ) - setattr( - self, - f"hc_{which}_scale_", - ParameterWeight( - weight_name=f"{self.prefix}.hc_{which}_scale", data_type=torch.float32, weight_shape=(3,) - ), - ) + p = self.prefix + self.hc_attn_fn_ = ParameterWeight( + weight_name=f"{p}.hc_attn_fn", + data_type=torch.float32, + weight_shape=(self.mix_hc, self.hc_mult * self.hidden), + ) + self.hc_attn_base_ = ParameterWeight( + weight_name=f"{p}.hc_attn_base", data_type=torch.float32, weight_shape=(self.mix_hc,) + ) + self.hc_attn_scale_ = ParameterWeight( + weight_name=f"{p}.hc_attn_scale", data_type=torch.float32, weight_shape=(3,) + ) + self.hc_ffn_fn_ = ParameterWeight( + weight_name=f"{p}.hc_ffn_fn", + data_type=torch.float32, + weight_shape=(self.mix_hc, self.hc_mult * self.hidden), + ) + self.hc_ffn_base_ = ParameterWeight( + weight_name=f"{p}.hc_ffn_base", data_type=torch.float32, weight_shape=(self.mix_hc,) + ) + self.hc_ffn_scale_ = ParameterWeight( + weight_name=f"{p}.hc_ffn_scale", data_type=torch.float32, weight_shape=(3,) + ) # ------------------------------------------------------------------ loading def load_hf_weights(self, weights): self._dequant_in_place(weights) return super().load_hf_weights(weights) - def _direct_fp8_weight_names(self): - names = set() - for attr_name in dir(self): - attr = getattr(self, attr_name, None) - quant_method = getattr(attr, "quant_method", None) - if getattr(quant_method, "method_name", None) == "deepgemm-fp8w8a8-b128": - names.update(getattr(attr, "weight_names", [])) - return names + def _fp8_scale_renames(self): + """Map weight name -> the scale name its quant method loads (e.g. `weight_scale_inv` + for DeepGEMM). Read from each MM weight's own `weight_scale_names`, so the rename + target always matches what that weight will look up; no-quant weights have None + entries and are skipped.""" + renames = {} + for attr in self.__dict__.values(): + weight_names = getattr(attr, "weight_names", ()) + scale_names = getattr(attr, "weight_scale_names", ()) + for weight_name, scale_name in zip(weight_names, scale_names): + if scale_name is not None: + renames[weight_name] = scale_name + return renames def _dequant_in_place(self, weights): p = self.prefix + "." - direct_fp8_names = self._direct_fp8_weight_names() - # Convert every (weight, scale) pair belonging to this layer. Existing FP8 matmul - # weights stay quantized; bmm-only weights are expanded; routed FP4 experts stay packed. - for k in [k for k in list(weights.keys()) if k.startswith(p) and k.endswith(".weight")]: - scale_k = k[: -len(".weight")] + ".scale" - if scale_k not in weights: - continue - w, s = weights[k], weights[scale_k] - if w.dtype == torch.int8: # FP4 routed experts stay packed for DeepseekV4FP4ExpertsWeight. + scale_renames = self._fp8_scale_renames() + # Convert every `.scale` belonging to this layer. Weights are loaded incrementally + # per safetensors shard, so the paired weight may live in another shard: + # - routed FP4 experts keep `.scale` as-is (matches marlin-mxfp4w4a16-b32's suffix); + # - FP8 matmul scales only need renaming for DeepGEMM, no weight required; + # - FP8 pairs on no-quant paths (wo_a's ROWBMMWeight) are expanded to bf16, + # the only case that truly requires weight and scale in the same shard. + for scale_k in [k for k in list(weights.keys()) if k.startswith(p) and k.endswith(".scale")]: + if scale_k.startswith(f"{p}ffn.experts."): continue - elif k in direct_fp8_names: # FP8 e4m3, block-128 scale, run by DeepGEMM directly - weights[k.replace("weight", "weight_scale_inv")] = s.to(torch.float32) + k = scale_k[: -len(".scale")] + ".weight" + target = scale_renames.get(k) + if target is not None: # FP8 e4m3, block-128 scale, run by DeepGEMM directly + weights[target] = weights[scale_k].to(torch.float32) del weights[scale_k] - else: # FP8 e4m3 for no-quant paths such as ROWBMMWeight - weights[k] = dequant_fp8_block_to_bf16(w, s).to(self.data_type_) + else: + weights[k] = dequant_fp8_block_to_bf16(weights[k], weights[scale_k]).to(self.data_type_) del weights[scale_k] # grouped-O: reshape [groups*o_lora, in] -> [groups, in, o_lora] for the batched matmul woa = f"{self.prefix}.attn.wo_a.weight" diff --git a/lightllm/models/deepseek_v4/mem_manager.py b/lightllm/models/deepseek_v4/mem_manager.py deleted file mode 100644 index 288d433380..0000000000 --- a/lightllm/models/deepseek_v4/mem_manager.py +++ /dev/null @@ -1,12 +0,0 @@ -from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import Deepseek2MemoryManager - - -class DeepseekV4MemoryManager(Deepseek2MemoryManager): - """Stores the per-token MLA KV (head_num=1, head_dim=512), reusing the deepseek2 layout/operator. - - The prefill path computes attention in-layer from the request's hidden states, so it does not read - this buffer. The decode/incremental path (M6) will add the sliding-window ring + compressed-KV + - per-request compressor-state buffers here. - """ - - pass diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 687d5f46f0..1a88e08977 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -7,11 +7,6 @@ from lightllm.models.llama.model import LlamaTpPartModel from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.basemodel.attention.base_att import ( - BaseAttBackend, - BasePrefillAttState, - BaseDecodeAttState, -) from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import ( DeepseekV4PreAndPostLayerWeight, ) @@ -27,12 +22,13 @@ from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import ( DeepseekV4TransformerLayerInfer, ) +from lightllm.common.basemodel.attention.create_utils import nsa_data_type_to_backend from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, linear_ramp_mask, ) -from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num +from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num, get_env_start_args from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager @@ -40,36 +36,6 @@ DSV4_DECODE_CUDAGRAPH_MAX_LEN = 8192 -class DeepseekV4DirectSparseAttBackend(BaseAttBackend): - """Lifecycle placeholder for V4 direct attention. - - V4 attention is currently driven inside the layer, not by the generic - `infer_state.prefill_att_state.prefill_att()` / `decode_att()` backend selector. - """ - - def create_att_prefill_state(self, infer_state: DeepseekV4InferStateInfo): - return DeepseekV4DirectSparsePrefillAttState(backend=self, infer_state=infer_state) - - def create_att_decode_state(self, infer_state: DeepseekV4InferStateInfo): - return DeepseekV4DirectSparseDecodeAttState(backend=self, infer_state=infer_state) - - -class DeepseekV4DirectSparsePrefillAttState(BasePrefillAttState): - def init_state(self): - return - - def prefill_att(self, *args, **kwargs): - raise RuntimeError("DeepSeek-V4 attention is executed directly in layer_infer.") - - -class DeepseekV4DirectSparseDecodeAttState(BaseDecodeAttState): - def init_state(self): - return - - def decode_att(self, *args, **kwargs): - raise RuntimeError("DeepSeek-V4 attention is executed directly in layer_infer.") - - @ModelRegistry("deepseek_v4") class DeepseekV4TpPartModel(LlamaTpPartModel): req_manager: DeepseekV4ReqManager @@ -107,6 +73,7 @@ def _init_req_manager(self): compress_rates=self._dsv4_compress_rates, head_dim=self.config["head_dim"], indexer_head_dim=self.config["index_head_dim"], + sliding_window=self.config["sliding_window"], ) return @@ -117,6 +84,10 @@ def _get_compress_rates(self, layer_num): def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() compress_rates = getattr(self, "_dsv4_compress_rates", self._get_compress_rates(layer_num)) + sliding_window = int(self.config["sliding_window"]) + # 活跃窗口之外的 swa 余量: 在途 prefill chunk 的瞬时占用(出窗槽位到下一次 prep 才回收) + # + radix cache 持有的窗口尾部(每条缓存序列约一个 window)。 + swa_extra_token_num = int(self.batch_max_tokens or 0) + self.max_req_num * sliding_window self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, @@ -126,7 +97,8 @@ def _init_mem_manager(self): compress_rates=compress_rates, indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, - sliding_window=self.config["sliding_window"], + sliding_window=sliding_window, + swa_extra_token_num=swa_extra_token_num, mem_fraction=self.mem_fraction, ) assert isinstance(self.req_manager, DeepseekV4ReqManager) @@ -137,7 +109,7 @@ def _init_cudagraph(self): if not self.disable_cudagraph and self.graph_max_len_in_batch > DSV4_DECODE_CUDAGRAPH_MAX_LEN: logger.info( "DeepSeek-V4 caps decode cudagraph max_len_in_batch from %s to %s for the current " - "graph-safe sparse-attention fallback; longer decode batches run eager.", + "graph-safe sparse-attention path; longer decode batches run eager.", self.graph_max_len_in_batch, DSV4_DECODE_CUDAGRAPH_MAX_LEN, ) @@ -150,8 +122,14 @@ def _can_run_prefill_cudagraph(self, infer_state: DeepseekV4InferStateInfo, hand return False def _init_att_backend(self): - self.prefill_att_backend = DeepseekV4DirectSparseAttBackend(model=self) - self.decode_att_backend = DeepseekV4DirectSparseAttBackend(model=self) + args = get_env_start_args() + if args.llm_kv_type == "None": + args.llm_kv_type = "fp8kv_dsa" + if args.llm_kv_type != "fp8kv_dsa": + raise RuntimeError("DeepSeek-V4 requires llm_kv_type=fp8kv_dsa for packed FlashMLA sparse attention") + backend_cls = nsa_data_type_to_backend["fp8kv_dsa"]["flashmla_sparse"] + self.prefill_att_backend = backend_cls(model=self) + self.decode_att_backend = backend_cls(model=self) return def _init_custom(self): @@ -165,11 +143,12 @@ def _init_custom(self): return def _init_to_get_rotary(self): - # Interleaved (GPT-J) rope. Build real cos/sin tables (_cos_cached_*/_sin_cached_*) following the - # gemma4 two-variant convention; the infer-struct slices them into position_cos_*/position_sin_* - # and apply_rotary_emb (interleaved, NOT the NeoX rotary_emb_fwd) applies them. Sliding-window - # layers use base rope_theta (no YaRN); compressed (CSA/HCA) layers use compress_rope_theta with - # YaRN. Tables kept fp32 for accuracy (the apply upcasts anyway). + # Interleaved (GPT-J) rope. Build complex64 freqs_cis tables (_freqs_cis_*) following the + # gemma4 two-variant convention; the fused sglang q kernel consumes them directly, while + # _cos_cached_*/_sin_cached_* are .real/.imag views of the same storage for the kv rope, + # inverse rope and compressor paths (apply_rotary_emb: interleaved, NOT the NeoX + # rotary_emb_fwd). Sliding-window layers use base rope_theta (no YaRN); compressed (CSA/HCA) + # layers use compress_rope_theta with YaRN. Kept fp32 for accuracy (the apply upcasts anyway). cfg = self.config rs = cfg.get("rope_scaling", {}) or {} dim = cfg["qk_rope_head_dim"] @@ -185,18 +164,29 @@ def build(base, factor, orig_max): smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() freqs = freqs / factor * (1 - smooth) + freqs * smooth f = torch.outer(torch.arange(max_seq, dtype=torch.float32, device="cuda"), freqs) # [max_seq, dim//2] - return f.cos(), f.sin() + return torch.complex(f.cos(), f.sin()) - self._cos_cached_sliding, self._sin_cached_sliding = build( + self._freqs_cis_sliding = build( cfg["rope_theta"], rs.get("factor", 16), rs.get("original_max_position_embeddings", 65536), ) - self._cos_cached_compress, self._sin_cached_compress = build( + self._freqs_cis_compress = build( cfg["compress_rope_theta"], rs.get("factor", 16), rs.get("original_max_position_embeddings", 65536), ) + self._cos_cached_sliding = self._freqs_cis_sliding.real + self._sin_cached_sliding = self._freqs_cis_sliding.imag + self._cos_cached_compress = self._freqs_cis_compress.real + self._sin_cached_compress = self._freqs_cis_compress.imag + # Each layer uses exactly one rope variant; wire its table once here (layers are already + # built: _init_infer_layer runs before _init_custom) instead of relaying via infer_state. + # The compressor needs the full compress tables (entry rope positions != token positions). + for layer in self.layers_infer: + layer.freqs_cis = self._freqs_cis_compress if layer.compress_ratio else self._freqs_cis_sliding + layer.cos_compress_table = self._cos_cached_compress + layer.sin_compress_table = self._sin_cached_compress return diff --git a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py new file mode 100644 index 0000000000..3510b92c30 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py @@ -0,0 +1,92 @@ +import torch + +import triton +import triton.language as tl + + +@triton.jit +def _fwd_kernel_destindex_copy_indexer_k_dsv4( + K, + Dest_loc, + O_fp8, + O_f32, + stride_k_bs, + stride_k_d, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, + SCALE_MIN: tl.constexpr, + HEAD_DIM: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BYTES_PER_PAGE: tl.constexpr, +): + cur_index = tl.program_id(0) + dest_index = tl.load(Dest_loc + cur_index).to(tl.int64) + # negative dest (unmapped slot) is a no-op, not an OOB write into a neighboring page. + if dest_index < 0: + return + + page = dest_index // PAGE_SIZE + token_in_page = dest_index % PAGE_SIZE + + offs_d = tl.arange(0, HEAD_DIM) + vals = tl.load(K + cur_index * stride_k_bs + offs_d * stride_k_d).to(tl.float32) + amax = tl.max(tl.abs(vals), axis=0) + # per-token plain fp32 scale (not ue8m0), matching DeepseekV4MemoryManager._pack_indexer_k + scale = tl.maximum(amax / FP8_MAX, SCALE_MIN) + k_fp8 = tl.clamp(vals / scale, min=FP8_MIN, max=FP8_MAX).to(tl.float8e4nv) + + data_base = page * BYTES_PER_PAGE + token_in_page * HEAD_DIM + tl.store(O_fp8 + data_base + offs_d, k_fp8) + scale_idx = (page * BYTES_PER_PAGE + PAGE_SIZE * HEAD_DIM) // 4 + token_in_page + tl.store(O_f32 + scale_idx, scale) + return + + +@torch.no_grad() +def destindex_copy_indexer_k_dsv4( + K: torch.Tensor, + DestLoc: torch.Tensor, + O_buffer: torch.Tensor, + page_size: int, +): + """Packed indexer-K page-slab writer (DeepSeek-V4 c4/CSA layers). + + K: [T, 128] bf16 unquantized indexer keys. + DestLoc: [T] int — c4-pool-local token slots; must already be allocated by the caller. + Negative slots (unmapped) are skipped. + O_buffer: [num_pages, bytes_per_page] uint8 — one layer's slab from the c4 indexer + PackedPagePool (128B fp8 data region + 4B fp32 scale tail per token). + + Bit-compatible with DeepseekV4MemoryManager._pack_indexer_k + PackedPagePool.write. + """ + seq_len = DestLoc.shape[0] + if seq_len == 0: + return + head_dim, scale_bytes = 128, 4 + + K = K.reshape(-1, head_dim) + assert K.shape[0] == seq_len, f"Expected K shape[0]={seq_len}, got {K.shape[0]}" + assert K.dtype == torch.bfloat16, f"Expected bf16 indexer K, got {K.dtype}" + bytes_per_page = O_buffer.shape[-1] + assert O_buffer.dtype == torch.uint8 and O_buffer.is_contiguous() + assert bytes_per_page % 4 == 0 + assert bytes_per_page >= page_size * (head_dim + scale_bytes) + + flat = O_buffer.view(-1) + _fwd_kernel_destindex_copy_indexer_k_dsv4[(seq_len,)]( + K, + DestLoc, + flat.view(torch.float8_e4m3fn), + flat.view(torch.float32), + K.stride(0), + K.stride(1), + FP8_MIN=torch.finfo(torch.float8_e4m3fn).min, + FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, + SCALE_MIN=1e-4, + HEAD_DIM=head_dim, + PAGE_SIZE=page_size, + BYTES_PER_PAGE=bytes_per_page, + num_warps=1, + num_stages=1, + ) + return diff --git a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py new file mode 100644 index 0000000000..a3ec6ed8cf --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py @@ -0,0 +1,121 @@ +import torch + +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def _fwd_kernel_destindex_copy_kv_flashmla_dsv4( + KV, + Dest_loc, + O_fp8, + O_bf16, + O_u8, + stride_kv_bs, + stride_kv_d, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, + SCALE_MIN: tl.constexpr, + NOPE_DIM: tl.constexpr, + ROPE_DIM: tl.constexpr, + GROUP_SIZE: tl.constexpr, + NUM_GROUPS: tl.constexpr, + SCALE_BYTES: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BYTES_PER_PAGE: tl.constexpr, +): + cur_index = tl.program_id(0) + dest_index = tl.load(Dest_loc + cur_index).to(tl.int64) + # negative dest (unmapped slot, e.g. full_to_c* rows that never closed a group) is a no-op, + # not an OOB write into a neighboring page. + if dest_index < 0: + return + + page = dest_index // PAGE_SIZE + token_in_page = dest_index % PAGE_SIZE + data_base = page * BYTES_PER_PAGE + token_in_page * (NOPE_DIM + ROPE_DIM * 2) + scale_base = page * BYTES_PER_PAGE + PAGE_SIZE * (NOPE_DIM + ROPE_DIM * 2) + token_in_page * SCALE_BYTES + + # nope: per-group ue8m0 quant. SCALE_BYTES(=NUM_GROUPS+1) lanes cover the exponent bytes + # plus the trailing zero pad byte in one store. libdevice.log2 (not tl.log2, which is the + # approx instruction) and the bit-packed 2**e keep this bit-exact with the torch oracle + # DeepseekV4MemoryManager._pack_mla_kv. + offs_g = tl.arange(0, SCALE_BYTES) + offs_e = tl.arange(0, GROUP_SIZE) + group_mask = offs_g < NUM_GROUPS + kv_ptrs = KV + cur_index * stride_kv_bs + (offs_g[:, None] * GROUP_SIZE + offs_e[None, :]) * stride_kv_d + vals = tl.load(kv_ptrs, mask=group_mask[:, None], other=0.0).to(tl.float32) + amax = tl.max(tl.abs(vals), axis=1) + scale_exp = tl.ceil(libdevice.log2(tl.maximum(amax / FP8_MAX, SCALE_MIN))).to(tl.int32) + scale = ((scale_exp + 127) << 23).to(tl.float32, bitcast=True) + kv_fp8 = tl.clamp(vals / scale[:, None], min=FP8_MIN, max=FP8_MAX).to(tl.float8e4nv) + tl.store(O_fp8 + data_base + offs_g[:, None] * GROUP_SIZE + offs_e[None, :], kv_fp8, mask=group_mask[:, None]) + scale_bytes = tl.where(group_mask, scale_exp + 127, 0).to(tl.uint8) + tl.store(O_u8 + scale_base + offs_g, scale_bytes) + + # rope: bf16 passthrough into the data region right after the nope bytes + offs_r = tl.arange(0, ROPE_DIM) + rope = tl.load(KV + cur_index * stride_kv_bs + (NOPE_DIM + offs_r) * stride_kv_d) + tl.store(O_bf16 + (data_base + NOPE_DIM) // 2 + offs_r, rope) + return + + +@torch.no_grad() +def destindex_copy_kv_flashmla_dsv4( + KV: torch.Tensor, + DestLoc: torch.Tensor, + O_buffer: torch.Tensor, + page_size: int, +): + """fp8_ds_mla packed page-slab writer (DeepSeek-V4 ABI, all latent pools). + + KV: [T, 512] bf16 — 448 normed-latent dims + 64 rope'd dims per token. + DestLoc: [T] int — pool-local token slots (page = slot // page_size); the pool HOLD slot is + a valid in-bounds row, negative slots (unmapped) are skipped. Slots must already be + resolved/allocated by the caller. + O_buffer: [num_pages, bytes_per_page] uint8 — one layer's slab from PackedPagePool + (swa page=128 / c4 page=64 / c128 page=2 all share this kernel). + + Per token: 448B fp8(e4m3) in 7x64 ue8m0 groups + 128B bf16 rope in the page data region; + 7 exponent bytes (e+127) + 1 zero pad at the page scale tail. Bit-compatible with + DeepseekV4MemoryManager._pack_mla_kv + PackedPagePool.write. + """ + seq_len = DestLoc.shape[0] + if seq_len == 0: + return + nope_dim, rope_dim, group_size = 448, 64, 64 + head_dim = nope_dim + rope_dim + scale_bytes = nope_dim // group_size + 1 + + KV = KV.reshape(-1, head_dim) + assert KV.shape[0] == seq_len, f"Expected KV shape[0]={seq_len}, got {KV.shape[0]}" + assert KV.dtype == torch.bfloat16, f"Expected bf16 KV (rope bytes are stored as-is), got {KV.dtype}" + bytes_per_page = O_buffer.shape[-1] + assert O_buffer.dtype == torch.uint8 and O_buffer.is_contiguous() + assert bytes_per_page % 2 == 0 + assert bytes_per_page >= page_size * (nope_dim + rope_dim * 2 + scale_bytes) + + flat = O_buffer.view(-1) + _fwd_kernel_destindex_copy_kv_flashmla_dsv4[(seq_len,)]( + KV, + DestLoc, + flat.view(torch.float8_e4m3fn), + flat.view(torch.bfloat16), + flat, + KV.stride(0), + KV.stride(1), + FP8_MIN=torch.finfo(torch.float8_e4m3fn).min, + FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, + SCALE_MIN=1e-4, + NOPE_DIM=nope_dim, + ROPE_DIM=rope_dim, + GROUP_SIZE=group_size, + NUM_GROUPS=nope_dim // group_size, + SCALE_BYTES=scale_bytes, + PAGE_SIZE=page_size, + BYTES_PER_PAGE=bytes_per_page, + num_warps=4, + num_stages=1, + ) + return diff --git a/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py index c7d2d59ec6..47d87d4932 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py +++ b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py @@ -1,15 +1,5 @@ import torch -# DeepSeek-V4-Flash ships weights in two quantized formats: -# * non-expert linears: FP8 e4m3 with block-[128,128] scales stored as float8_e8m0fnu (ue8m0) -# * routed experts: FP4 e2m1 packed 2-per-byte (stored as int8) with group-32 ue8m0 scales -# Hopper (H200) has no native SM100 MegaMoE path. Non-expert FP8 weights can run directly through -# DeepGEMM. Routed FP4 experts are converted blockwise to FP8, avoiding a full bf16 expansion. - -# OCP E2M1 magnitude table for the 3 low bits (sign = bit 3). torch.float4_e2m1fn_x2 packs two -# such codes per byte, low nibble = lower (even) logical index. -_E2M1_MAG = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0] - def e8m0_to_fp32(scale: torch.Tensor) -> torch.Tensor: """float8_e8m0fnu encodes 2**(byte-127); torch decodes it correctly on .to(float32).""" @@ -24,70 +14,3 @@ def dequant_fp8_block_to_bf16(weight_e4m3: torch.Tensor, scale_e8m0: torch.Tenso s = e8m0_to_fp32(scale_e8m0).cuda().contiguous() # weight_dequant runs with torch default dtype for the output; force bf16 result. return weight_dequant(w, s, block_size) - - -def cast_e2m1fn_to_e4m3fn(weight_int8: torch.Tensor, scale_e8m0: torch.Tensor): - """Cast packed FP4 e2m1 expert weights to FP8 e4m3 with block-128 fp32 scales. - - This follows the DeepSeek-V4 reference converter, but returns the scale in fp32 because - LightLLM's DeepGEMM FP8 weight pack stores block scales as fp32. - """ - assert weight_int8.dtype == torch.int8 - assert weight_int8.ndim == 2 - out_dim, packed_in = weight_int8.shape - in_dim = packed_in * 2 - fp8_block_size = 128 - fp4_block_size = 32 - assert in_dim % fp8_block_size == 0 and out_dim % fp8_block_size == 0 - assert scale_e8m0.shape[0] == out_dim - assert scale_e8m0.shape[1] == in_dim // fp4_block_size - - table = torch.tensor( - [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0], - dtype=torch.float32, - device=weight_int8.device, - ) - packed = weight_int8.view(torch.uint8) - low = packed & 0x0F - high = (packed >> 4) & 0x0F - vals = torch.stack([table[low.long()], table[high.long()]], dim=-1).reshape(out_dim, in_dim) - - # 6.0 * 2**6 fits in e4m3fn (384 < 448), while 6.0 * 2**7 would overflow. - max_offset_bits = 6 - block_out = out_dim // fp8_block_size - block_in = in_dim // fp8_block_size - - vals = vals.view(block_out, fp8_block_size, block_in, fp8_block_size).transpose(1, 2) - scale = scale_e8m0.float().view(block_out, fp8_block_size, block_in, -1).transpose(1, 2).flatten(2) - block_scale = scale.amax(dim=-1, keepdim=True) / (2**max_offset_bits) - offset = scale / block_scale - offset = offset.unflatten(-1, (fp8_block_size, -1)).repeat_interleave(fp4_block_size, dim=-1) - vals = (vals * offset).transpose(1, 2).reshape(out_dim, in_dim) - block_scale = block_scale.squeeze(-1).to(torch.float8_e8m0fnu).to(torch.float32) - return vals.to(torch.float8_e4m3fn), block_scale - - -def dequant_fp4_group_to_bf16(weight_int8: torch.Tensor, scale_e8m0: torch.Tensor, group_size: int = 32): - """De-quantize an int8-packed FP4 e2m1 weight to bf16. - - weight_int8: [out, in // 2] int8 (two e2m1 codes per byte, low nibble = even index). - scale_e8m0: [out, in // group_size] ue8m0 (one scale per group_size logical elements along K). - returns: [out, in] bf16. - """ - w = weight_int8.cuda() - out, packed_in = w.shape - in_dim = packed_in * 2 - b = w.to(torch.int32).bitwise_and(0xFF) - lut = torch.tensor(_E2M1_MAG, dtype=torch.float32, device=w.device) - - def _decode(nib: torch.Tensor) -> torch.Tensor: - mag = lut[nib.bitwise_and(0x7)] - neg = nib.bitwise_and(0x8).bool() - return torch.where(neg, -mag, mag) - - lo = _decode(b.bitwise_and(0xF)) - hi = _decode(b.bitwise_right_shift(4).bitwise_and(0xF)) - vals = torch.stack([lo, hi], dim=-1).reshape(out, in_dim) # [out, in] - s = e8m0_to_fp32(scale_e8m0).cuda() # [out, in//group_size] - s = s.repeat_interleave(group_size, dim=1)[:, :in_dim] - return (vals * s).to(torch.bfloat16) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 2db6c67e77..70d5c72ac3 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -625,8 +625,11 @@ def make_argument_parser() -> argparse.ArgumentParser: type=str, default=None, choices=["fp8", "fp4"], - help="""Expert quantization dtype for EP MoE. Supported values are - fp8 and fp4. Note that fp4 is only supported on SM100 GPUs.""", + help="""Requested dtype for MoE expert weights, fp8 or fp4. Resolves the fused_moe + quant method: fp8 -> deepgemm-fp8w8a8-b128; fp4 -> deepgemm-fp4fp8-b32 (online + quantization) on SM100 GPUs, or marlin-mxfp4w4a16-b32 (Marlin W4A16, TP only) on other GPUs. + Defaults to `expert_dtype` in config.json if present. Per-layer override: + --quant_cfg mix_bits with name `fused_moe`.""", ) parser.add_argument( "--vit_quant_type", diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index 8be5198eb3..ff09c018be 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -316,6 +316,9 @@ def _insert_helper_no_recursion( def match_prefix(self, key, update_refs=False): key = key[: self._align_len(len(key))] + if len(key) == 0: + return None, 0, None + key = self._trim_key_by_extra_value_validity(key) if len(key) == 0: return None, 0, None ans_value_list = [] @@ -331,6 +334,30 @@ def match_prefix(self, key, update_refs=False): self.dec_node_ref_counter(self.root_node) return None, 0, None + def _trim_key_by_extra_value_validity(self, key: torch.Tensor) -> torch.Tensor: + """命中有效性裁剪(extra_value_ops 提供 valid_match_length 时启用,如 DeepSeek-V4 的 + swa 按页 bitmap): 先做一次只读探测遍历得到自然命中与沿路 extra_value,按其有效边界截短 + key,随后的正常遍历(加引用/分裂)只走截短后的前缀 —— 引用计数与最终返回值在同一次遍历 + 内保持一致,不存在事后裁剪导致的漏减/多减。 + + 探测遍历可能分裂部分命中的节点(与正常遍历同语义,树不变式不受影响)。裁剪只会缩短命中, + 没有任何失败路径。""" + if self.extra_value_ops is None: + return key + valid_match_length = getattr(self.extra_value_ops, "valid_match_length", None) + if valid_match_length is None: + return key + probe_values = [] + probe_node = self._match_prefix_helper(self.root_node, key, probe_values, update_refs=False) + if probe_node == self.root_node or len(probe_values) == 0: + return key + natural_len = sum(len(v) for v in probe_values) + extra_value = self.get_extra_value_by_node(probe_node) + valid_len = int(valid_match_length(extra_value, natural_len)) + if valid_len < natural_len: + return key[:valid_len] + return key + def _match_prefix_helper( self, node: TreeNode, key: torch.Tensor, ans_value_list: list, update_refs=False ) -> TreeNode: @@ -595,6 +622,36 @@ def _print_helper(self, node: TreeNode, indent): self._print_helper(child, indent=indent + 2) return + def reclaim_unreferenced_swa_pages(self, need_pages: int) -> None: + """DeepSeek-V4 swa 压力阀: 页 allocator 触底时,沿 LRU 序(evict_tree_set)只对 + ref_count==0 的节点链回收其 swa 页(full 槽与压缩条目保留——节点仍可服务更长前缀的 + 中段命中),并清载荷 bitmap 位使后续命中按缩短语义裁剪。所有权判定直接复用 radix + 引用计数: 节点被任何活跃请求借用即 ref>0,其页不可达。不够时由 allocator 的 assert + 兜底(最后防线)。""" + if self.mem_manager is None or self.extra_value_ops is None: + return + invalidate = getattr(self.extra_value_ops, "invalidate_swa_pages", None) + if invalidate is None: + return + allocator = self.mem_manager.swa_page_allocator + target = allocator.can_use_mem_size + int(need_pages) + for leaf in list(self.evict_tree_set): + if allocator.can_use_mem_size >= target: + break + node = leaf + # 叶子起步沿父链回收: 引用计数向上累加(add_node_ref_counter 走父链), + # 因此 ref==0 的祖先必无任何活跃借用方。重复访问无害(evict_swa/-1 跳过)。 + # 每回收一个节点就复查目标,避免多回收(无谓削减命中可用性)。 + while node is not None and node is not self.root_node and node.ref_counter == 0: + if len(node.token_mem_index_value) > 0: + self.mem_manager.evict_swa(node.token_mem_index_value) + if node.token_extra_value is not None: + invalidate(node.token_extra_value) + if allocator.can_use_mem_size >= target: + return + node = node.parent + return + def free_radix_cache_to_get_enough_token(self, need_token_num): assert self.mem_manager is not None if need_token_num > self.mem_manager.allocator.can_use_mem_size: diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 3fd6e0463a..32b7f6b3d5 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -124,40 +124,19 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: return req_objs - def free_a_req_mem( - self, - free_token_index: List, - req: "InferReq", - free_c4_index: Optional[List] = None, - free_c128_index: Optional[List] = None, - ): + def free_a_req_mem(self, free_token_index: List, req: "InferReq"): is_dsv4_req_manager = hasattr(self.req_manager, "build_prompt_cache_payload") - if hasattr(self.req_manager, "pop_compress_indices_for_req") and not is_dsv4_req_manager: - c4, c128 = self.req_manager.pop_compress_indices_for_req(req.req_idx) - if c4 is not None and free_c4_index is not None: - free_c4_index.append(c4) - if c128 is not None and free_c128_index is not None: - free_c128_index.append(c128) - self.req_manager.clear_runtime_state(req.req_idx) - if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) if is_dsv4_req_manager: - c4, c128 = self.req_manager.pop_compress_indices_for_req(req.req_idx) - if c4 is not None and free_c4_index is not None: - free_c4_index.append(c4) - if c128 is not None and free_c128_index is not None: - free_c128_index.append(c128) - self.req_manager.clear_runtime_state(req.req_idx) + # 槽位随 full 槽经 mem_manager.free 级联回收。pause 路径不释放 req_idx, + # 必须在此复位出窗水位线 + 清 c128 在途状态(恢复命中走 extend,不会再有 + # restore/zero 时机;c4 状态随 swa 页生灭,无需处理)。 + self.req_manager.init_compress_state(req.req_idx) else: if not self.is_linear_att_mixed_model: if is_dsv4_req_manager: - self._dsv4_full_att_free_req( - free_token_index=free_token_index, - req=req, - free_c4_index=free_c4_index, - free_c128_index=free_c128_index, - ) + self._dsv4_full_att_free_req(free_token_index=free_token_index, req=req) else: self._full_att_free_req(free_token_index=free_token_index, req=req) else: @@ -187,79 +166,60 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): req.shared_kv_node = None return - def _dsv4_full_att_free_req( - self, - free_token_index: List, - req: "InferReq", - free_c4_index: Optional[List] = None, - free_c128_index: Optional[List] = None, - ): + def _dsv4_full_att_free_req(self, free_token_index: List, req: "InferReq"): if req.cur_kv_len == 0: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0:0]) return old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len - cache_len = self.radix_cache.align_len(req.cur_kv_len) inserted_len = old_prefix_len duplicate_prefix_len = old_prefix_len - inserted_payload = None - pending_payload = getattr(req, "prompt_cache_snapshot_payload", None) - pending_cache_len = getattr(req, "prompt_cache_snapshot_len", 0) - - # The current V4 runtime state is only guaranteed to describe the current - # sequence end. Cache aligned current ends; leave unaligned tails uncached. - if pending_payload is not None and pending_cache_len > old_prefix_len: - cache_len = pending_cache_len - input_token_ids = req.get_input_token_ids() - key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") - value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() - duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value, extra_value=pending_payload) - inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len - if inserted_len == cache_len: - inserted_payload = pending_payload - else: - self.req_manager.release_prompt_cache_detached_swa(pending_payload) - pending_payload = None - inserted_len = old_prefix_len - duplicate_prefix_len = old_prefix_len - elif cache_len == req.cur_kv_len and cache_len > old_prefix_len: - input_token_ids = req.get_input_token_ids() - key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") - value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() + + # 载荷只剩按页 bitmap(compressor 状态随 swa 页生灭/边界自然归零,不进载荷), + # 任意 128 对齐前缀皆可插入——含生成段(floor(cur_kv_len) 边界,回收保留尾页保证其驻留)。 + cache_len = self.radix_cache.align_len(req.cur_kv_len) + if cache_len > old_prefix_len: payload = self.req_manager.build_prompt_cache_payload(req.req_idx, cache_len) - duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value, extra_value=payload) - inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len - if inserted_len == cache_len: - inserted_payload = payload - self.req_manager.detach_prompt_cache_payload_from_req(req.req_idx, inserted_payload) - else: - inserted_len = old_prefix_len - duplicate_prefix_len = old_prefix_len + value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() + # 按页有效性 bitmap 用插入时刻的映射写定(此后只会被阀清 0,不会复活)。水位线 + # 纯 CPU 推导,避免 router 关键路径上的 GPU gather 同步(每插入一次要等全部在途 + # decode kernel)。插入门: 截掉结尾的 invalid 页 —— 它们生来不可命中,还会永久 + # 挡住后续更长前缀复用同一段 token(全量重插会因前缀已存在而保留旧 bitmap)。 + page_size = self.req_manager.get_prompt_cache_page_size() + bitmap = self.req_manager.swa_page_valid_from_watermark(req.req_idx, cache_len) + n_pages = int(bitmap.numel()) + while n_pages > 0 and not bool(bitmap[n_pages - 1]): + n_pages -= 1 + gated_len = n_pages * page_size + if gated_len < cache_len: + logger.info( + f"DeepSeek-V4 prompt cache insert gate: trailing swa pages already evicted, " + f"shrink insert {cache_len} -> {gated_len}" + ) + cache_len = gated_len + payload.cache_len = cache_len + payload.swa_page_valid = bitmap[:n_pages].clone() + + if cache_len > old_prefix_len: + input_token_ids = req.get_input_token_ids() + key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") + duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value[:cache_len], extra_value=payload) + inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len + if inserted_len != cache_len: + inserted_len = old_prefix_len + duplicate_prefix_len = old_prefix_len - if ( - pending_payload is not None - and inserted_payload is not pending_payload - and pending_cache_len <= old_prefix_len - ): - self.req_manager.release_prompt_cache_detached_swa(pending_payload) - req.prompt_cache_snapshot_payload = None - req.prompt_cache_snapshot_len = 0 dense_row = self.req_manager.req_to_token_indexs[req.req_idx] self._append_free_token_index(free_token_index, dense_row[old_prefix_len:duplicate_prefix_len]) self._append_free_token_index(free_token_index, dense_row[inserted_len : req.cur_kv_len]) if len(free_token_index) == 0: free_token_index.append(dense_row[0:0]) + # 释放的 full 槽经 mem_manager.free 级联回收 swa/c4/c128(映射键控,无需收集槽位)。 - c4, c128 = self.req_manager.pop_prompt_cache_free_compress_indices( - req.req_idx, - keep_len=inserted_len, - duplicate_start_len=old_prefix_len, - duplicate_end_len=duplicate_prefix_len, - ) - if c4 is not None and free_c4_index is not None: - free_c4_index.append(c4) - if c128 is not None and free_c128_index is not None: - free_c128_index.append(c128) + # pause 路径不会走 req_manager.free/init: 复位出窗水位线(残留水位线会破坏下一次 + # prefill 的共享前缀保护)并清 c128 在途状态(恢复命中走 extend 续算,若残留暂停前的 + # 半窗聚合会算错;c128 状态在 128 对齐命中边界本应为零)。 + self.req_manager.init_compress_state(req.req_idx) if req.shared_kv_node is not None: assert req.shared_kv_node.node_prefix_total_len <= max(inserted_len, old_prefix_len) @@ -375,13 +335,11 @@ def _filter(self, finished_request_ids: List[int]): free_req_index = [] free_token_index = [] - free_c4_index = [] - free_c128_index = [] for request_id in finished_request_ids: req: InferReq = self.requests_mapping.pop(request_id) if self.args.diverse_mode: req.clear_master_slave_state() - self.free_a_req_mem(free_token_index, req, free_c4_index, free_c128_index) + self.free_a_req_mem(free_token_index, req) free_req_index.append(req.req_idx) # logger.info(f"infer release req id {req.shm_req.request_id}") @@ -389,17 +347,7 @@ def _filter(self, finished_request_ids: List[int]): self.shm_req_manager.put_back_req_obj(req.shm_req) free_token_index = custom_cat(free_token_index) - if hasattr(self.req_manager, "free_compress_indices"): - free_c4_index = custom_cat(free_c4_index) if free_c4_index else None - free_c128_index = custom_cat(free_c128_index) if free_c128_index else None - self.req_manager.free( - free_req_index, - free_token_index, - free_c4_index=free_c4_index, - free_c128_index=free_c128_index, - ) - else: - self.req_manager.free(free_req_index, free_token_index) + self.req_manager.free(free_req_index, free_token_index) finished_req_ids_set = set(finished_request_ids) self.infer_req_ids = [_id for _id in self.infer_req_ids if _id not in finished_req_ids_set] @@ -428,13 +376,11 @@ def pause_reqs(self, pause_reqs: List["InferReq"], is_master_in_dp: bool): g_infer_state_lock.acquire() free_token_index = [] - free_c4_index = [] - free_c128_index = [] for req in pause_reqs: if self.args.diverse_mode: # 发生暂停的时候,需要清除 diverse 模式下的主从关系 req.clear_master_slave_state() - self.free_a_req_mem(free_token_index, req, free_c4_index, free_c128_index) + self.free_a_req_mem(free_token_index, req) assert req.wait_pause is True req.wait_pause = False req.paused = True @@ -445,13 +391,6 @@ def pause_reqs(self, pause_reqs: List["InferReq"], is_master_in_dp: bool): if len(free_token_index) != 0: free_token_index = custom_cat(free_token_index) self.req_manager.free_token(free_token_index) - if hasattr(self.req_manager, "free_compress_indices"): - free_c4_index = custom_cat(free_c4_index) if free_c4_index else None - free_c128_index = custom_cat(free_c128_index) if free_c128_index else None - self.req_manager.free_compress_indices( - free_c4_index=free_c4_index, - free_c128_index=free_c128_index, - ) g_infer_state_lock.release() return self @@ -738,8 +677,6 @@ def _init_all_state(self): g_infer_context.req_manager.req_sampling_params_manager.init_req_sampling_params(self) if hasattr(g_infer_context.req_manager, "init_compress_state"): g_infer_context.req_manager.init_compress_state(req_idx=self.req_idx) - self.prompt_cache_snapshot_len = 0 - self.prompt_cache_snapshot_payload = None self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() # token healing mode 才被使用的管理对象 @@ -775,11 +712,9 @@ def _match_radix_cache(self): ready_cache_len = share_node.node_prefix_total_len # 从 cpu 到 gpu 是流内阻塞操作 g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor - if hasattr(g_infer_context.req_manager, "restore_prompt_cache_payload"): - payload = g_infer_context.radix_cache.get_extra_value_by_node(share_node) - if payload is None: - raise RuntimeError("DeepSeek-V4 radix cache hit is missing prompt-cache payload") - g_infer_context.req_manager.restore_prompt_cache_payload(self.req_idx, payload) + # DeepSeek-V4 命中无需任何恢复: 槽位由 full_to_* 映射键控(radix 持有 full 槽即有效, + # 命中长度已在 match_prefix 内按 bitmap 裁剪),c4 compressor 状态随 swa 页常驻 + # (零拷贝续算),c128 状态在 128 对齐边界自然归零(init_compress_state 已清)。 self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 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 74fdb1e87b..4c5e1af222 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -212,6 +212,9 @@ def init_model(self, kvargs): page_size=radix_page_size, extra_value_ops=radix_extra_value_ops, ) + if radix_extra_value_ops is not None and hasattr(self.model.mem_manager, "set_swa_pressure_valve"): + # swa 页 allocator 触底时让 radix 对 ref==0 节点回收 swa 页(DeepSeek-V4)。 + self.model.mem_manager.set_swa_pressure_valve(self.radix_cache.reclaim_unreferenced_swa_pages) if "prompt_cache_kv_buffer" in model_cfg: assert self.use_dynamic_prompt_cache @@ -713,31 +716,6 @@ def _pre_handle_finished_reqs(self, finished_reqs: List[InferReq]): """ pass - def _maybe_capture_prompt_cache_payload(self, req_obj: InferReq): - if self.radix_cache is None: - return - req_manager = g_infer_context.req_manager - if not hasattr(req_manager, "build_prompt_cache_payload"): - return - if req_obj.sampling_param.disable_prompt_cache: - return - page_size = getattr(self.args, "dynamic_prompt_cache_page_size", 1) - cache_len = int(req_obj.cur_kv_len) - if page_size <= 1 or cache_len <= 0 or cache_len % page_size != 0: - return - if cache_len > req_obj.shm_req.input_len: - return - if getattr(req_obj, "prompt_cache_snapshot_len", 0) >= cache_len: - return - - payload = req_manager.build_prompt_cache_payload(req_obj.req_idx, cache_len, clone_swa=True) - old_payload = getattr(req_obj, "prompt_cache_snapshot_payload", None) - if old_payload is not None: - req_manager.release_prompt_cache_detached_swa(old_payload, keep_payload=payload) - req_obj.prompt_cache_snapshot_len = cache_len - req_obj.prompt_cache_snapshot_payload = payload - return - # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] @@ -785,7 +763,6 @@ def _post_handle( ): req_obj: InferReq = req_obj pack: InferReqUpdatePack = pack - self._maybe_capture_prompt_cache_payload(req_obj) pack.handle( next_token_id=next_token_id, next_token_logprob=next_token_logprob, From b3b81237cd886fbe697c8835b24aedbd55db4078 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 11 Jun 2026 07:31:59 +0000 Subject: [PATCH 009/214] fix --- .../deepseek4_mem_manager.py | 8 - .../layer_infer/hyper_connection.py | 4 +- .../layer_infer/post_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 139 +++++++++++------- 4 files changed, 87 insertions(+), 65 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 47fdf76fd9..3df5c0e12e 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -34,7 +34,6 @@ # c4 compressor state ring(overlap 对: 每页 2 个分组槽 × ratio 4 行)。c128 state 在 128 边界 # 自然归零(在线聚合),无缓存常驻需求,保持 req 键控,不进 swa 派生池。 DSV4_C4_STATE_RING = 8 -DSV4_PROFILE_MAX_FULL_TOKENS = 1_500_000 # swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 的对应物)。 # lightllm 的调度准入只看 full 池,prefill 优先的波次会让"已 prefill 未 decode"的请求整段 # prompt 占住 swa 槽(首次 decode prep 才批量出窗回收),峰值≈准入波次 prompt 总和。在 @@ -276,13 +275,6 @@ def profile_size(self, mem_fraction): dist.all_reduce(tensor, op=dist.ReduceOp.MIN) self.size = tensor.item() - if self.size > DSV4_PROFILE_MAX_FULL_TOKENS: - logger.info( - f"DeepseekV4MemoryManager cap profiled max_total_token_num from " - f"{self.size} to {DSV4_PROFILE_MAX_FULL_TOKENS} to keep runtime headroom" - ) - self.size = DSV4_PROFILE_MAX_FULL_TOKENS - logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" f"{str(self.get_cell_size() / 1024 ** 2)} MB is the conservative size of one token kv cache\n" diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index b125e9ed06..080ebabd89 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -58,9 +58,9 @@ def hc_post(x, residual, post_mix, res_mix): return torch.ops.vllm.mhc_post_tilelang(x, residual, post_mix, res_mix) -def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, rms_eps, hc_eps): +def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, rms_eps, hc_eps, alloc_func): """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim].""" - out = torch.empty(streams.shape[0], dim, device=streams.device, dtype=streams.dtype) + out = alloc_func((streams.shape[0], dim), dtype=streams.dtype, device=streams.device) torch.ops.vllm.hc_head_fused_kernel_tilelang( streams.view(-1, hc_mult, dim).contiguous(), hc_fn, diff --git a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py index bc95c249f7..8eddfb3b9d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py @@ -22,5 +22,6 @@ def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, l cfg["hidden_size"], cfg["rms_norm_eps"], cfg.get("hc_eps", 1e-6), + self.alloc_tensor, ) return super().token_forward(collapsed, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 6f7c0de3fc..95119586fd 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -4,6 +4,7 @@ from lightllm.common.basemodel import TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.distributed.communication_op import all_reduce +from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post @@ -66,7 +67,7 @@ def __init__(self, layer_num, network_config): self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 # ------------------------------------------------------------------ forward (HC-threaded) - def _hc_attn_in(self, input_embdings, layer_weight): + def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): """Layer input -> attention input (attn_norm fused). First layer gets the raw streams and runs a standalone hc_pre; later layers get (x, residual, post_mix, res_mix) and fuse the previous layer's ffn hc_post with this layer's attn hc_pre.""" @@ -99,7 +100,7 @@ def _hc_attn_in(self, input_embdings, layer_weight): self.eps_, ) - def _hc_ffn_in(self, x, residual, post_mix, res_mix, layer_weight): + def _hc_ffn_in(self, x, residual, post_mix, res_mix, layer_weight: DeepseekV4TransformerLayerWeight): """Attention output -> ffn input (ffn_norm fused): fused attn hc_post + ffn hc_pre.""" return hc_fused_post_pre( x, @@ -124,14 +125,18 @@ def _hc_ffn_out(self, x, residual, post_mix, res_mix): streams = hc_post(x, residual, post_mix, res_mix) return streams.reshape(streams.shape[0], -1) - def context_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): + def context_forward( + self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): x, residual, post_mix, res_mix = self._hc_attn_in(input_embdings, layer_weight) x = self.context_attention_forward(x, infer_state, layer_weight) x, residual, post_mix, res_mix = self._hc_ffn_in(x, residual, post_mix, res_mix, layer_weight) x = self._ffn(x, infer_state, layer_weight) return self._hc_ffn_out(x, residual, post_mix, res_mix) - def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight): + def token_forward( + self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): x, residual, post_mix, res_mix = self._hc_attn_in(input_embdings, layer_weight) x = self.token_attention_forward(x, infer_state, layer_weight) x, residual, post_mix, res_mix = self._hc_ffn_in(x, residual, post_mix, res_mix, layer_weight) @@ -144,36 +149,39 @@ def _select_rope(self, infer_state: DeepseekV4InferStateInfo): return infer_state.position_cos_compress, infer_state.position_sin_compress return infer_state.position_cos_sliding, infer_state.position_sin_sliding - def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): from sglang.jit_kernel.dsv4 import fused_q_norm_rope + x = self._tpsp_allgather(input=x, infer_state=infer_state) cos_tok, sin_tok = self._select_rope(infer_state) T = x.shape[0] qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(x), eps=self.eps_) q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim) # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, # fused in one sglang dsv4 jit kernel (fp32 norm/rotation, bf16 in between -- same as eager). - q = torch.empty_like(q_in) + q = self.alloc_tensor(q_in.shape, dtype=q_in.dtype, device=q_in.device) fused_q_norm_rope(q_in, q, self.eps_, self.freqs_cis, infer_state.position_ids) - kv = layer_weight.kv_norm_(layer_weight.wkv_.mm(x), eps=self.eps_) - kv = torch.cat( - [ - kv[:, : -self.rope_dim], - apply_rotary_emb(kv[:, -self.rope_dim :], cos_tok, sin_tok), - ], - dim=1, + # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 sglang jit kernel 完成 + # (同 sglang _compute_kv_to_cache),替代 eager norm/rope/cat + _post_cache_kv。 + # bf16 kv 中间量没有其他消费者: flashmla 路径注意力读 cache,压缩器/indexer 取 x。 + infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( + layer_index=self.layer_num_, + mem_index=infer_state.mem_index, + kv=layer_weight.wkv_.mm(x), + kv_weight=layer_weight.kv_norm_.weight, + eps=self.eps_, + freqs_cis=self.freqs_cis, + positions=infer_state.position_ids, ) - return q, kv, qa, cos_tok, sin_tok + return q, qa, cos_tok, sin_tok - def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): # o: [T, tp_q_heads, head_dim] after inverse rope -> grouped low-rank O -> [T, hidden] T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] o = layer_weight.wo_b_.mm(o) - if self.tp_world_size_ > 1: - all_reduce(o, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - return o + return self._tpsp_reduce(input=o, infer_state=infer_state) def _inv_rope(self, o, cos_tok, sin_tok): return torch.cat( @@ -190,7 +198,9 @@ def _inv_rope(self, o, cos_tok, sin_tok): ) # ------------------------------------------------------------------ compressor / indexer - def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _indexer_q_weight( + self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): if self.compress_ratio != 4: return None, None cos_tok = infer_state.position_cos_compress @@ -222,7 +232,7 @@ def _write_compressed_kv(self, infer_state: DeepseekV4InferStateInfo, req, entry infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) return slots - def _compressor_weights(self, layer_weight, for_indexer: bool): + def _compressor_weights(self, layer_weight: DeepseekV4TransformerLayerWeight, for_indexer: bool): if for_indexer: return ( layer_weight.idx_cmp_wkv_.mm_param.weight, @@ -239,7 +249,9 @@ def _compressor_weights(self, layer_weight, for_indexer: bool): self.head_dim, ) - def _run_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _run_compressor_prefill( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): """Per-request compressor for the prefill chunk. Runs as part of the deferred attention func, before the attention metadata gathers the slot mappings. @@ -255,7 +267,9 @@ def _run_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, laye self._run_c128_compressor_prefill(x, infer_state, layer_weight) return - def _run_c4_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _run_c4_compressor_prefill( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): rm = infer_state.req_manager mem = infer_state.mem_manager wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) @@ -316,7 +330,9 @@ def _run_c4_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, l infer_state.mem_manager.pack_indexer_k_to_cache(self.layer_num_, slots, idx_comp) return - def _run_c128_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _run_c128_compressor_prefill( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): rm = infer_state.req_manager wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) b_req = infer_state.b_req_idx.tolist() @@ -365,7 +381,9 @@ def _run_c128_compressor_prefill(self, x, infer_state: DeepseekV4InferStateInfo, self._write_compressed_kv(infer_state, req, entry_start, entry.unsqueeze(0)) return - def _run_compressor_decode(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _run_compressor_decode( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): """Batched decode compressor (cuda-graph safe): state update for every request, cache write masked to the pool HOLD slot unless this token completes a window. Compressed-cache slots were pre-allocated by prepare_decode_compress_slots in the prep phase. @@ -460,24 +478,24 @@ def _run_compressor_decode(self, x, infer_state: DeepseekV4InferStateInfo, layer return # ------------------------------------------------------------------ attention (prefill) - def context_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): - q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) - # template hook: write the chunk's packed latent into the swa pool before attention - # reads it back via full_to_swa indices (this custom forward bypasses the tpl path). - self._post_cache_kv(cache_kv, infer_state, layer_weight) - o = self._context_attention_wrapper_run(q, cache_kv, q_lora, x, infer_state, layer_weight) + def context_attention_forward( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): + # _get_qkv writes the chunk's packed latent into the swa pool (fused kernel) before + # attention reads it back via full_to_swa indices (this custom forward bypasses the + # tpl _post_cache_kv path). + q, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + o = self._context_attention_wrapper_run(q, q_lora, x, infer_state, layer_weight) return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) def _context_attention_wrapper_run( - self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight + self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): if torch.cuda.is_current_stream_capturing(): q = q.contiguous() - cache_kv = cache_kv.contiguous() q_lora = q_lora.contiguous() x = x.contiguous() _q = tensor_to_no_ref_tensor(q) - _cache_kv = tensor_to_no_ref_tensor(cache_kv) _q_lora = tensor_to_no_ref_tensor(q_lora) _x = tensor_to_no_ref_tensor(x) @@ -486,11 +504,13 @@ def _context_attention_wrapper_run( infer_state.prefill_cuda_graph_create_graph_obj() infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() - o = torch.empty((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) + # Same graph-split output handoff as the template, but avoid its dry-run because + # DSV4 attention mutates compressor/cache state before returning. + o = self.alloc_tensor((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) _o = tensor_to_no_ref_tensor(o) def att_func(new_infer_state: DeepseekV4InferStateInfo): - tmp_o = self._context_attention_kernel(_q, _cache_kv, _q_lora, _x, new_infer_state, layer_weight) + tmp_o = self._context_attention_kernel(_q, _q_lora, _x, new_infer_state, layer_weight) assert tmp_o.shape == _o.shape _o.copy_(tmp_o) return @@ -498,9 +518,11 @@ def att_func(new_infer_state: DeepseekV4InferStateInfo): infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) return o - return self._context_attention_kernel(q, cache_kv, q_lora, x, infer_state, layer_weight) + return self._context_attention_kernel(q, q_lora, x, infer_state, layer_weight) - def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _context_attention_kernel( + self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): self._run_compressor_prefill(x, infer_state, layer_weight) idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) att_control = AttControl( @@ -511,7 +533,6 @@ def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: Deepsee "compress_ratio": self.compress_ratio, "head_dim_v": self.head_dim, "softmax_scale": self.softmax_scale, - "cache_kv": cache_kv, "q_lora": q_lora, "hidden_states": x, "attn_sink": layer_weight.attn_sink_.weight, @@ -530,13 +551,16 @@ def _context_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: Deepsee ) # ------------------------------------------------------------------ attention (decode) - def token_attention_forward(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): - q, cache_kv, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) - self._post_cache_kv(cache_kv, infer_state, layer_weight) - o = self._token_attention_kernel(q, cache_kv, q_lora, x, infer_state, layer_weight) + def token_attention_forward( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): + q, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + o = self._token_attention_kernel(q, q_lora, x, infer_state, layer_weight) return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) - def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _token_attention_kernel( + self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): self._run_compressor_decode(x, infer_state, layer_weight) idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) att_control = AttControl( @@ -547,7 +571,6 @@ def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV "compress_ratio": self.compress_ratio, "head_dim_v": self.head_dim, "softmax_scale": self.softmax_scale, - "cache_kv": cache_kv, "q_lora": q_lora, "hidden_states": x, "attn_sink": layer_weight.attn_sink_.weight, @@ -566,7 +589,7 @@ def _token_attention_kernel(self, q, cache_kv, q_lora, x, infer_state: DeepseekV ) # ------------------------------------------------------------------ moe - def _routed_experts(self, x, weights, indices, layer_weight): + def _routed_experts(self, x, weights, indices, layer_weight: DeepseekV4TransformerLayerWeight): return layer_weight.experts_.experts_with_preselected( input_tensor=x, topk_weights=weights, @@ -574,7 +597,11 @@ def _routed_experts(self, x, weights, indices, layer_weight): clamp_limit=float(self.swiglu_limit), ) - def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): + x = x.view(-1, self.hidden) + if not self.enable_ep_moe: + x = self._tpsp_allgather(input=x, infer_state=infer_state) + gw = layer_weight.gate_weight_.mm_param.weight logits = F.linear(x.float(), gw.float()).contiguous() weights, indices = self._select_experts(logits, infer_state, layer_weight) @@ -582,7 +609,7 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): g = layer_weight.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) u = layer_weight.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) shared = layer_weight.shared_down_.mm((F.silu(g) * u).to(x.dtype)) - if self.enable_ep_moe and getattr(layer_weight.experts_, "is_ep", False): + if self.enable_ep_moe: if self.tp_world_size_ > 1: all_reduce( shared, @@ -592,14 +619,16 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight): ) return routed + shared out = routed + shared - if self.tp_world_size_ > 1: - all_reduce(out, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - return out + return self._tpsp_reduce(input=out, infer_state=infer_state) - def _select_experts(self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _select_experts( + self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): return self._select_experts_vllm(logits, infer_state, layer_weight) - def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _select_experts_vllm( + self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): from vllm import _custom_ops as ops M = logits.shape[0] @@ -616,9 +645,9 @@ def _select_experts_vllm(self, logits, infer_state: DeepseekV4InferStateInfo, la else: bias = layer_weight.gate_bias_.weight - weights = torch.empty((M, self.topk), dtype=torch.float32, device=logits.device) - indices = torch.empty((M, self.topk), dtype=indices_dtype, device=logits.device) - token_expert_indices = torch.empty((M, self.topk), dtype=torch.int32, device=logits.device) + weights = self.alloc_tensor((M, self.topk), dtype=torch.float32, device=logits.device) + indices = self.alloc_tensor((M, self.topk), dtype=indices_dtype, device=logits.device) + token_expert_indices = self.alloc_tensor((M, self.topk), dtype=torch.int32, device=logits.device) ops.topk_hash_softplus_sqrt( weights, indices, From 6002866d2fc7c7858c26f7383364eb75bf19edd8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 11 Jun 2026 12:47:19 +0000 Subject: [PATCH 010/214] fix rope --- .../deepseek2/triton_kernel/rotary_emb.py | 66 ++++++++++++------- .../layer_infer/transformer_layer_infer.py | 25 ++----- lightllm/models/deepseek_v4/model.py | 58 ++++++++-------- .../deepseek_v4/triton_kernel/rotary_emb.py | 26 -------- 4 files changed, 75 insertions(+), 100 deletions(-) delete mode 100644 lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py diff --git a/lightllm/models/deepseek2/triton_kernel/rotary_emb.py b/lightllm/models/deepseek2/triton_kernel/rotary_emb.py index 30e5a59248..a8f851de2a 100644 --- a/lightllm/models/deepseek2/triton_kernel/rotary_emb.py +++ b/lightllm/models/deepseek2/triton_kernel/rotary_emb.py @@ -29,6 +29,8 @@ def _rotary_kernel( BLOCK_SEQ: tl.constexpr, BLOCK_DMODEL: tl.constexpr, NUM_STAGE: tl.constexpr, + HAS_K: tl.constexpr, + INVERSE: tl.constexpr, ): head_start_index = tl.program_id(0) seq_block_index = tl.program_id(1) @@ -44,6 +46,8 @@ def _rotary_kernel( off_dimcos_sin = seq_index * stride_cosbs + cos_range * stride_cosd cos = tl.load(Cos + off_dimcos_sin) sin = tl.load(Sin + off_dimcos_sin) + if INVERSE: + sin = -sin if HEAD_PARALLEL_NUM == 1: for q_head_index in tl.static_range(0, HEAD_Q, step=1): @@ -56,18 +60,19 @@ def _rotary_kernel( tl.store(Q + off_q0, out_q0) tl.store(Q + off_q1, out_q1) - for k_head_index in tl.static_range(0, HEAD_K, step=1): - off_k0 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range0 * stride_kd - off_k1 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range1 * stride_kd + if HAS_K: + for k_head_index in tl.static_range(0, HEAD_K, step=1): + off_k0 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range0 * stride_kd + off_k1 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range1 * stride_kd - k0 = tl.load(K + off_k0) - k1 = tl.load(K + off_k1) + k0 = tl.load(K + off_k0) + k1 = tl.load(K + off_k1) - out_k0 = k0 * cos - k1 * sin - out_k1 = k0 * sin + k1 * cos + out_k0 = k0 * cos - k1 * sin + out_k1 = k0 * sin + k1 * cos - tl.store(K + off_k0, out_k0) - tl.store(K + off_k1, out_k1) + tl.store(K + off_k0, out_k0) + tl.store(K + off_k1, out_k1) else: for q_head_index in tl.range(head_start_index, HEAD_Q, step=HEAD_PARALLEL_NUM, num_stages=NUM_STAGE): off_q0 = seq_index * stride_qbs + q_head_index * stride_qh + dim_range0 * stride_qd @@ -79,18 +84,19 @@ def _rotary_kernel( tl.store(Q + off_q0, out_q0) tl.store(Q + off_q1, out_q1) - for k_head_index in tl.range(head_start_index, HEAD_K, step=HEAD_PARALLEL_NUM, num_stages=NUM_STAGE): - off_k0 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range0 * stride_kd - off_k1 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range1 * stride_kd + if HAS_K: + for k_head_index in tl.range(head_start_index, HEAD_K, step=HEAD_PARALLEL_NUM, num_stages=NUM_STAGE): + off_k0 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range0 * stride_kd + off_k1 = seq_index * stride_kbs + k_head_index * stride_kh + dim_range1 * stride_kd - k0 = tl.load(K + off_k0) - k1 = tl.load(K + off_k1) + k0 = tl.load(K + off_k0) + k1 = tl.load(K + off_k1) - out_k0 = k0 * cos - k1 * sin - out_k1 = k0 * sin + k1 * cos + out_k0 = k0 * cos - k1 * sin + out_k1 = k0 * sin + k1 * cos - tl.store(K + off_k0, out_k0) - tl.store(K + off_k1, out_k1) + tl.store(K + off_k0, out_k0) + tl.store(K + off_k1, out_k1) return @@ -109,7 +115,10 @@ def get_test_configs(): def get_static_key(q, k): - head_num_q, head_num_k, head_dim = q.shape[1], k.shape[1], q.shape[2] + assert q is not None, "q can not be None" + head_num_q = q.shape[1] + head_num_k = k.shape[1] if k is not None else 0 + head_dim = q.shape[2] return { "Q_HEAD_NUM": head_num_q, "K_HEAD_NUM": head_num_k, @@ -126,12 +135,17 @@ def get_static_key(q, k): mutates_args=["q", "k"], ) @torch.no_grad() -def rotary_emb_fwd(q, k, cos, sin, run_config=None): +def rotary_emb_fwd(q, k, cos, sin, inverse=False, run_config=None): + assert q is not None, "q can not be None" + has_k = k is not None and k.shape[1] != 0 total_len = q.shape[0] - head_num_q, head_num_k = q.shape[1], k.shape[1] + head_num_q = q.shape[1] + head_num_k = k.shape[1] if k is not None else 0 head_dim = q.shape[2] assert q.shape[0] == cos.shape[0] and q.shape[0] == sin.shape[0], f"q shape {q.shape} cos shape {cos.shape}" - assert k.shape[0] == cos.shape[0] and k.shape[0] == sin.shape[0], f"k shape {k.shape} cos shape {cos.shape}" + if k is not None: + assert k.shape[0] == cos.shape[0] and k.shape[0] == sin.shape[0], f"k shape {k.shape} cos shape {cos.shape}" + assert k.shape[2] == head_dim, f"k shape {k.shape} q head_dim {head_dim}" assert triton.next_power_of_2(head_dim) == head_dim if not run_config: @@ -157,9 +171,9 @@ def rotary_emb_fwd(q, k, cos, sin, run_config=None): stride_qbs=q.stride(0), stride_qh=q.stride(1), stride_qd=q.stride(2), - stride_kbs=k.stride(0), - stride_kh=k.stride(1), - stride_kd=k.stride(2), + stride_kbs=k.stride(0) if k is not None else 0, + stride_kh=k.stride(1) if k is not None else 0, + stride_kd=k.stride(2) if k is not None else 0, stride_cosbs=cos.stride(0), stride_cosd=cos.stride(1), stride_sinbs=sin.stride(0), @@ -171,6 +185,8 @@ def rotary_emb_fwd(q, k, cos, sin, run_config=None): BLOCK_SEQ=BLOCK_SEQ, BLOCK_DMODEL=head_dim, NUM_STAGE=num_stages, + HAS_K=has_k, + INVERSE=inverse, num_warps=num_warps, num_stages=num_stages, ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 95119586fd..e610473170 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -17,7 +17,7 @@ paged_prefill_compress_data, paged_decode_state_slots, ) -from ..triton_kernel.rotary_emb import apply_rotary_emb +from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo @@ -184,18 +184,9 @@ def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: Deepsee return self._tpsp_reduce(input=o, infer_state=infer_state) def _inv_rope(self, o, cos_tok, sin_tok): - return torch.cat( - [ - o[..., : -self.rope_dim], - apply_rotary_emb( - o[..., -self.rope_dim :], - cos_tok.unsqueeze(1), - sin_tok.unsqueeze(1), - inverse=True, - ), - ], - dim=-1, - ) + # in-place; 单张量路径只需要旋转 rope 切片。 + rotary_emb_fwd(o[..., -self.rope_dim :], None, cos_tok, sin_tok, inverse=True) + return o # ------------------------------------------------------------------ compressor / indexer def _indexer_q_weight( @@ -206,13 +197,7 @@ def _indexer_q_weight( cos_tok = infer_state.position_cos_compress sin_tok = infer_state.position_sin_compress idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_heads, self.index_head_dim) - idx_q = torch.cat( - [ - idx_q[..., : -self.rope_dim], - apply_rotary_emb(idx_q[..., -self.rope_dim :], cos_tok.unsqueeze(1), sin_tok.unsqueeze(1)), - ], - dim=-1, - ) + rotary_emb_fwd(idx_q[..., -self.rope_dim :], None, cos_tok, sin_tok) idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale return idx_q, idx_weight diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 1a88e08977..63430e548b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -22,7 +22,7 @@ from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import ( DeepseekV4TransformerLayerInfer, ) -from lightllm.common.basemodel.attention.create_utils import nsa_data_type_to_backend +from lightllm.common.basemodel.attention import get_nsa_prefill_att_backend_class, get_nsa_decode_att_backend_class from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, @@ -125,11 +125,11 @@ def _init_att_backend(self): args = get_env_start_args() if args.llm_kv_type == "None": args.llm_kv_type = "fp8kv_dsa" + # TODO: 支持其他 kv type if args.llm_kv_type != "fp8kv_dsa": raise RuntimeError("DeepSeek-V4 requires llm_kv_type=fp8kv_dsa for packed FlashMLA sparse attention") - backend_cls = nsa_data_type_to_backend["fp8kv_dsa"]["flashmla_sparse"] - self.prefill_att_backend = backend_cls(model=self) - self.decode_att_backend = backend_cls(model=self) + self.prefill_att_backend = get_nsa_prefill_att_backend_class(index=0)(model=self) + self.decode_att_backend = get_nsa_decode_att_backend_class(index=0)(model=self) return def _init_custom(self): @@ -146,36 +146,36 @@ def _init_to_get_rotary(self): # Interleaved (GPT-J) rope. Build complex64 freqs_cis tables (_freqs_cis_*) following the # gemma4 two-variant convention; the fused sglang q kernel consumes them directly, while # _cos_cached_*/_sin_cached_* are .real/.imag views of the same storage for the kv rope, - # inverse rope and compressor paths (apply_rotary_emb: interleaved, NOT the NeoX - # rotary_emb_fwd). Sliding-window layers use base rope_theta (no YaRN); compressed (CSA/HCA) - # layers use compress_rope_theta with YaRN. Kept fp32 for accuracy (the apply upcasts anyway). + # inverse rope and compressor paths (deepseek2's interleaved triton rotary_emb_fwd). + # Sliding-window layers use base rope_theta (no YaRN); + # compressed (CSA/HCA) layers use compress_rope_theta with configured rope_scaling. + # Kept fp32 for accuracy (the apply upcasts anyway). cfg = self.config rs = cfg.get("rope_scaling", {}) or {} dim = cfg["qk_rope_head_dim"] - beta_fast = rs.get("beta_fast", 32) - beta_slow = rs.get("beta_slow", 1) max_seq = max(int(self.max_seq_length), int(cfg.get("max_position_embeddings", 8192))) max_seq = min(max_seq, 1 << 18) # cap table size (256K) for correctness-first - - def build(base, factor, orig_max): - freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim)) - if orig_max > 0: - low, high = find_correction_range(beta_fast, beta_slow, dim, base, orig_max) - smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() - freqs = freqs / factor * (1 - smooth) + freqs * smooth - f = torch.outer(torch.arange(max_seq, dtype=torch.float32, device="cuda"), freqs) # [max_seq, dim//2] - return torch.complex(f.cos(), f.sin()) - - self._freqs_cis_sliding = build( - cfg["rope_theta"], - rs.get("factor", 16), - rs.get("original_max_position_embeddings", 65536), - ) - self._freqs_cis_compress = build( - cfg["compress_rope_theta"], - rs.get("factor", 16), - rs.get("original_max_position_embeddings", 65536), - ) + freq_exponents = torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim + positions = torch.arange(max_seq, dtype=torch.float32, device="cuda") + + sliding_freqs = 1.0 / (cfg["rope_theta"] ** freq_exponents) + f = torch.outer(positions, sliding_freqs) # [max_seq, dim//2] + self._freqs_cis_sliding = torch.complex(f.cos(), f.sin()) + + compress_freqs = 1.0 / (cfg["compress_rope_theta"] ** freq_exponents) + rope_type = rs.get("rope_type", rs.get("type", "default")) + orig_max = rs.get("original_max_position_embeddings", 0) + if rope_type == "yarn" and orig_max > 0: + beta_fast = rs.get("beta_fast", 32) + beta_slow = rs.get("beta_slow", 1) + factor = rs.get("factor", 1) + if factor is None: + factor = cfg.get("max_position_embeddings", max_seq) / orig_max + low, high = find_correction_range(beta_fast, beta_slow, dim, cfg["compress_rope_theta"], orig_max) + smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() + compress_freqs = compress_freqs / factor * (1 - smooth) + compress_freqs * smooth + f = torch.outer(positions, compress_freqs) # [max_seq, dim//2] + self._freqs_cis_compress = torch.complex(f.cos(), f.sin()) self._cos_cached_sliding = self._freqs_cis_sliding.real self._sin_cached_sliding = self._freqs_cis_sliding.imag self._cos_cached_compress = self._freqs_cis_compress.real diff --git a/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py b/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py deleted file mode 100644 index cb50977446..0000000000 --- a/lightllm/models/deepseek_v4/triton_kernel/rotary_emb.py +++ /dev/null @@ -1,26 +0,0 @@ -import torch - -# Interleaved (GPT-J) rotary application for DeepSeek-V4. Unlike llama/gemma's NeoX-style -# rotary_emb_fwd (rotate-half: pairs channel i with i+d/2 over a real cos/sin table), V4 rotates -# adjacent pairs (x0,x1),(x2,x3),... — a different channel pairing — so it cannot reuse -# rotary_emb_fwd, but it consumes the same real cos/sin tables (built in model.py:_init_to_get_rotary -# as _cos_cached_*/_sin_cached_*, gemma4-style). Correctness-first pure-torch; a fused triton port is -# a perf follow-up. - - -def apply_rotary_emb(x, cos, sin, inverse=False): - """Apply interleaved rope to the LAST dim of x (size = 2*cos.size(-1)). - - x: [..., rope_dim] (real). cos/sin: [..., rope_dim//2], broadcastable to x's paired view. - For x of shape [N, H, rope_dim], pass cos/sin [N, 1, rope_dim//2]; for [N, rope_dim] pass [N, rope_dim//2]. - Returns a new tensor of x's dtype (not in-place). inverse=True applies the conjugate rotation. - """ - dtype = x.dtype - x = x.float().reshape(*x.shape[:-1], -1, 2) - x0, x1 = x[..., 0], x[..., 1] - cos = cos.float() - sin = sin.float() - if inverse: - sin = -sin - out = torch.stack([x0 * cos - x1 * sin, x0 * sin + x1 * cos], dim=-1) - return out.flatten(-2).to(dtype) From 6bc34adb3896473a8eda46609c232b571faa7d8e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 11 Jun 2026 06:53:24 +0000 Subject: [PATCH 011/214] dsv4: enable decode cudagraph; fix warmup-baked FlashMLASchedMeta Root cause of the historical cudagraph accuracy drop (gsm8k 0.96 -> 0.74, coherent-but-runaway generations; same 0.75 the pre-v5 fullslot_decode experiments worked around): _capture_decode warms up via copy.copy(infer_state), which SHARES decode_att_state. FlashMLASchedMeta is lazily planned at the first kernel call and written back onto that shared state, so the warmup pass locks a schedule planned for the dummy batch (seq=2); the capture pass then binds those stale scheduler tensors and every replay runs real requests with a tile schedule planned for near-empty kv (systematically under-read attention). Fix: reset_sched_meta_for_capture() hook on the nsa decode att state, invoked in both capture paths after warmup, so planning happens INSIDE the captured region and re-plans on every replay from live tensors. Validation (tp4, H200, prompt cache on): batch-1 greedy decode is now character-identical to eager; per-layer probe shows embed+swa layers bitwise equal under replay, benign rounding-class deltas only in compress layers, argmax unchanged. gsm8k 100q/128: cold 0.960/111s, warm 0.960/23.3s 100% hits (eager: 0.95-0.97, cold 141s / warm 50s). Batch-1 decode 20.4ms/token vs 142ms eager. 41/41 unit tests green. Codex review GO (incl. overlap-path symmetry). launch.sh: drop --disable_cudagraph, derive PYTHONPATH from the script dir (hardcoded tree path made a worktree launch silently serve main-tree code). Co-Authored-By: Claude Fable 5 --- launch.sh | 39 +++++++++++++++++++ .../attention/nsa/fp8_flashmla_sparse.py | 8 ++++ lightllm/common/basemodel/cuda_graph.py | 19 +++++++++ 3 files changed, 66 insertions(+) create mode 100644 launch.sh diff --git a/launch.sh b/launch.sh new file mode 100644 index 0000000000..ce016efc96 --- /dev/null +++ b/launch.sh @@ -0,0 +1,39 @@ +# DeepSeek-V4-Flash serving (run inside the lightllm container, repo mounted at /data/wanzihao/lightllm-ds4). +# Verified 2026-06-11: smoke + gsm8k pass with this configuration (prompt cache ENABLED, decode +# cudagraph ENABLED; gsm8k 100q/128: cold 0.960/112s, warm 0.970/23.5s with 100% cache hits — +# vs eager cold 0.970/141s, warm 0.960/50s; batch-1 decode 20.4ms/token vs 142ms eager). +# +# Required env/flags and why: +# LOADWORKER=16 - parallel weight loading (~5x faster startup). +# Optional sizing knobs (defaults shown): LIGHTLLM_DSV4_SWA_FULL_TOKENS_RATIO=0.1 (swa pool floor +# as a fraction of full tokens; raise for long-prompt x high-parallel workloads), +# LIGHTLLM_DSV4_PROFILE_MAX_FULL_TOKENS=1500000 (auto-profile cap on max_total_token_num). +# PYTHONPATH sglang - _get_qkv / compressor reuse sglang.jit_kernel.dsv4 (fused_q_norm_rope, compress_old). +# --batch_max_tokens 8192 - FlashMLA get_decoding_sched_meta rejects >8192 rows per call (probed: 8192 OK, 12288 fails). +# decode cudagraph ENABLED - the v5 decode path is graph-safe: slot alloc/scatter in prep (outside +# graph), forward is pure gathers, HOLD padding rows redirect to HOLD slots. CORRECTNESS NOTE: +# FlashMLASchedMeta is lazily planned at first kernel call and written back onto the (shared) +# decode att state; the capture warmup pass would bake a dummy-content plan into the graph +# (gsm8k dropped to 0.74 with coherent-but-runaway generations). reset_sched_meta_for_capture() +# in cuda_graph._capture_decode re-plans INSIDE the captured region so every replay re-plans. +# DSV4 caps graph max_len_in_batch at 8192; longer decode batches fall back to eager. +# --disable_flashinfer_allreduce - flashinfer cuda_ipc resolves libcudart to tilelang's stub (undefined cudaDeviceReset); symm-mem allreduce is used instead. +# +# One-time container setup already applied (survives until container rebuild): +# pip install ipython (sglang import dependency) +# site-packages/vllm: layers/mhc.py + kernels/mhc/ + _tilelang_ops.py overlaid from /data/wanzihao/vllm (mhc_pre_tilelang ops; original kept at layers/mhc.py.bak) +# +# original: python -m lightllm.server.api_server --model_dir /data/models/DeepSeek-V4-Flash --tp 4 --enable_prefill_cudagraph + +# repo root = this script's directory, so the same file works in the main tree and in worktrees +# (a hardcoded tree path here once made a worktree launch silently serve main-tree code). +REPO_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +LOADWORKER=16 \ +PYTHONPATH="${REPO_DIR}":/data/wanzihao/sglang/python \ +python -m lightllm.server.api_server \ + --model_dir /data/models/DeepSeek-V4-Flash \ + --tp 4 \ + --batch_max_tokens 8192 \ + --disable_flashinfer_allreduce \ + --port 8000 diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 0570adea83..14b1b3307d 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -464,6 +464,14 @@ def init_state(self): self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} return + def reset_sched_meta_for_capture(self): + # cuda-graph capture hook: the warmup pass already locked/stored sched meta on this + # (shared) state object; reset so the capture pass re-plans INSIDE the graph and every + # replay re-plans from the live tensors instead of binding warmup leftovers. + flash_mla = self.backend.flash_mla() + self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} + return + def decode_att( self, q: Tuple[torch.Tensor, torch.Tensor], diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py index 782150661e..e1d96b744e 100644 --- a/lightllm/common/basemodel/cuda_graph.py +++ b/lightllm/common/basemodel/cuda_graph.py @@ -14,6 +14,20 @@ logger = init_logger(__name__) +def _reset_att_state_sched_meta(infer_state: InferStateInfo): + # capture 前调用: warmup 趟用 copy.copy 浅拷贝共享 decode_att_state,其内部惰性初始化的 + # 调度对象(如 FlashMLASchedMeta,首次内核调用时按当时数据规划并回写)会被 warmup 的 + # dummy 负载锁定;若不重置,捕获趟将绑定为 dummy 规划的调度张量,所有 replay 都用错误 + # 的 tile schedule(DSV4 实测 gsm8k 0.96 -> 0.74)。重置后规划发生在捕获区内,随 replay 重算。 + for att_state in (infer_state.decode_att_state, infer_state.decode_att_state1): + if att_state is None: + continue + reset_fn = getattr(att_state, "reset_sched_meta_for_capture", None) + if reset_fn is not None: + reset_fn() + return + + class CudaGraph: # CudaGraph forward pass for the decoding stage. @@ -94,6 +108,8 @@ def _capture_decode(self, decode_func, infer_state: InferStateInfo): if param_name not in pure_para_set: delattr(infer_state, param_name) + _reset_att_state_sched_meta(infer_state) + with torch.cuda.graph(graph_obj, pool=self.mempool): model_output = decode_func(infer_state) self.graph[batch_size] = (graph_obj, infer_state, model_output) @@ -128,6 +144,9 @@ def _capture_decode_overlap( if para_name not in pure_para_set1: delattr(infer_state1, para_name) + _reset_att_state_sched_meta(infer_state) + _reset_att_state_sched_meta(infer_state1) + with torch.cuda.graph(graph_obj, pool=self.mempool): model_output, model_output1 = decode_func(infer_state, infer_state1) self.graph[batch_size] = ( From e78e0d429dff4a58aad5cb3e9ac8158b215d870f Mon Sep 17 00:00:00 2001 From: wanzihao Date: Thu, 11 Jun 2026 08:32:14 +0000 Subject: [PATCH 012/214] dsv4: enable prefill cudagraph; zero pad-row attention output Graph-sandwich prefill (graphs capture dense ops only; attention/compressor run eagerly between segments) was already in-tree; enabling it exposed that HOLD-pad rows read the racing HOLD slot, making their hiddens nondeterministic and perturbing real rows via MoE expert batching (ulp-level, amplified ~1.9x/layer). Zero the pad rows' attention output. Residual greedy-trajectory divergence vs eager equals the fp4 marlin MoE kernel's own run-to-run reduction-order noise (eager-vs-eager control: 0/4 match), accepted statistically: gsm8k 100q cold 0.980/115.5s warm 0.960/25.9s (eager-baseline parity); batch-1 TTFT 1.86x at 46 tokens. --- launch.sh | 17 ++++++++++++++--- lightllm/models/deepseek_v4/infer_struct.py | 7 +++++++ .../layer_infer/transformer_layer_infer.py | 7 ++++++- 3 files changed, 27 insertions(+), 4 deletions(-) diff --git a/launch.sh b/launch.sh index ce016efc96..ccf870dc8c 100644 --- a/launch.sh +++ b/launch.sh @@ -5,9 +5,6 @@ # # Required env/flags and why: # LOADWORKER=16 - parallel weight loading (~5x faster startup). -# Optional sizing knobs (defaults shown): LIGHTLLM_DSV4_SWA_FULL_TOKENS_RATIO=0.1 (swa pool floor -# as a fraction of full tokens; raise for long-prompt x high-parallel workloads), -# LIGHTLLM_DSV4_PROFILE_MAX_FULL_TOKENS=1500000 (auto-profile cap on max_total_token_num). # PYTHONPATH sglang - _get_qkv / compressor reuse sglang.jit_kernel.dsv4 (fused_q_norm_rope, compress_old). # --batch_max_tokens 8192 - FlashMLA get_decoding_sched_meta rejects >8192 rows per call (probed: 8192 OK, 12288 fails). # decode cudagraph ENABLED - the v5 decode path is graph-safe: slot alloc/scatter in prep (outside @@ -17,6 +14,18 @@ # (gsm8k dropped to 0.74 with coherent-but-runaway generations). reset_sched_meta_for_capture() # in cuda_graph._capture_decode re-plans INSIDE the captured region so every replay re-plans. # DSV4 caps graph max_len_in_batch at 8192; longer decode batches fall back to eager. +# --enable_prefill_cudagraph + --prefill_cudagraph_max_handle_token 2048 - graph-sandwich prefill: +# graphs capture only the per-token dense ops; attention/compressor/indexer run eagerly between +# graph segments (att_func), so host-side planning and .tolist() prep never enter capture. Only +# cold prefills (prefix_total_token_num == 0, model gate) of <= 2048 new tokens replay; cache-hit +# and large batched prefills stay eager. Buckets are padded with a HOLD tail request whose +# attention output MUST be zeroed (infer_struct._dsv4_prefill_pad_q_len): pad rows read the +# racing HOLD slot, and nondeterministic pad hiddens perturb real rows via MoE expert batching +# (ulp-level, chaotically amplified ~1.9x/layer to O(1) by layer ~16 -> greedy token flips). +# Residual caveat: padded-vs-unpadded expert-batch composition still shifts reductions by ulps, +# same class as decode bucket padding; run-to-run determinism is anyway bounded by the fp4 +# marlin MoE kernel itself (probabilistic 1-ulp reduction-order noise measured eager-vs-eager). +# Acceptance is therefore statistical (gsm8k parity), not bitwise. # --disable_flashinfer_allreduce - flashinfer cuda_ipc resolves libcudart to tilelang's stub (undefined cudaDeviceReset); symm-mem allreduce is used instead. # # One-time container setup already applied (survives until container rebuild): @@ -36,4 +45,6 @@ python -m lightllm.server.api_server \ --tp 4 \ --batch_max_tokens 8192 \ --disable_flashinfer_allreduce \ + --enable_prefill_cudagraph \ + --prefill_cudagraph_max_handle_token 2048 \ --port 8000 diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index 39a6889d72..caf8ca6fa8 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -26,3 +26,10 @@ def init_some_extra_state(self, model): self.position_sin_sliding = torch.index_select(model._sin_cached_sliding, 0, pos) self.position_cos_compress = torch.index_select(model._cos_cached_compress, 0, pos) self.position_sin_compress = torch.index_select(model._sin_cached_compress, 0, pos) + # prefill-cudagraph 桶填充的 HOLD 尾请求的 q 行数。其注意力读 HOLD 槽位(内容被并发写 + # 竞争,每轮不同),输出必须清零,否则 pad 行 hidden 不确定 -> MoE 路由抖动 -> 共享 expert + # 批次组成变化 -> 真实行 GEMM 归约顺序变化(ulp 级),44 层放大后翻转低置信 token。 + self._dsv4_prefill_pad_q_len = 0 + if self.is_prefill and self.b_req_idx.numel() > 0: + if int(self.b_req_idx[-1].item()) == self.req_manager.HOLD_REQUEST_ID: + self._dsv4_prefill_pad_q_len = int((self.b_seq_len[-1] - self.b_ready_cache_len[-1]).item()) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index e610473170..761558c95d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -528,12 +528,17 @@ def _context_attention_kernel( "tp_world_size": self.tp_world_size_, }, ) - return infer_state.prefill_att_state.prefill_att( + out = infer_state.prefill_att_state.prefill_att( q=q, k=infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_), v=None, att_control=att_control, ) + pad_q_len = getattr(infer_state, "_dsv4_prefill_pad_q_len", 0) + if pad_q_len: + # pad 行读 HOLD 槽位(参见 infer_struct._dsv4_prefill_pad_q_len),清零以保持确定性 + out[-pad_q_len:] = 0 + return out # ------------------------------------------------------------------ attention (decode) def token_attention_forward( From c09dc6aa1d90a6de9b63135dd80748bbb2b7b27d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 11 Jun 2026 14:29:49 +0000 Subject: [PATCH 013/214] fix profile --- launch.sh | 3 + .../deepseek4_mem_manager.py | 43 ++++++++++-- lightllm/common/quantization/deepgemm.py | 70 +++++++++++++++++-- 3 files changed, 106 insertions(+), 10 deletions(-) diff --git a/launch.sh b/launch.sh index ccf870dc8c..b9c10d3f0a 100644 --- a/launch.sh +++ b/launch.sh @@ -7,6 +7,9 @@ # LOADWORKER=16 - parallel weight loading (~5x faster startup). # PYTHONPATH sglang - _get_qkv / compressor reuse sglang.jit_kernel.dsv4 (fused_q_norm_rope, compress_old). # --batch_max_tokens 8192 - FlashMLA get_decoding_sched_meta rejects >8192 rows per call (probed: 8192 OK, 12288 fails). +# kv pool sizing: auto-profiled from mem_fraction. The fp4 marlin MoE weights materialize their +# CUDA marlin-layout buffers at construction (MXFP4MoEQuantizationMethod._create_weight), so the +# profile sees the true weight footprint on any GPU/config. --max_total_token_num overrides. # decode cudagraph ENABLED - the v5 decode path is graph-safe: slot alloc/scatter in prep (outside # graph), forward is pure gathers, HOLD padding rows redirect to HOLD slots. CORRECTNESS NOTE: # FlashMLASchedMeta is lazily planned at first kernel call and written back onto the (shared) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 3df5c0e12e..f87a11704d 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -34,11 +34,10 @@ # c4 compressor state ring(overlap 对: 每页 2 个分组槽 × ratio 4 行)。c128 state 在 128 边界 # 自然归零(在线聚合),无缓存常驻需求,保持 req 键控,不进 swa 派生池。 DSV4_C4_STATE_RING = 8 -# swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 的对应物)。 -# lightllm 的调度准入只看 full 池,prefill 优先的波次会让"已 prefill 未 decode"的请求整段 -# prompt 占住 swa 槽(首次 decode prep 才批量出窗回收),峰值≈准入波次 prompt 总和。在 -# v5 的 swa 压力阀/准入耦合落地前,用比 sglang 更宽的 0.3 兜住该瞬时峰值。 -DSV4_SWA_FULL_TOKENS_RATIO = 0.3 +# swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 同值)。 +# v5 的 swa 压力阀(借页/驱逐)已覆盖 radix 树与准入波次的瞬时增长,结构性预算 +# (max_req×window + batch_max_tokens 余量)另行叠加,0.1 仅作 full 池比例下限。 +DSV4_SWA_FULL_TOKENS_RATIO = 0.1 def _ceil_div(a: int, b: int) -> int: @@ -702,6 +701,40 @@ def pack_mla_kv_to_cache(self, layer_index: int, mem_index: torch.Tensor, kv: to ) return + def pack_mla_kv_to_cache_fused_norm_rope( + self, + layer_index: int, + mem_index: torch.Tensor, + kv: torch.Tensor, + kv_weight: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, + ): + """同 pack_mla_kv_to_cache,但 rmsnorm + 尾部交错 rope 融合进写入 kernel + (sglang fused_k_norm_rope_flashmla,即 sglang _compute_kv_to_cache 的池侧), + 省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" + if kv.shape[0] == 0: + return + from sglang.jit_kernel.dsv4 import fused_k_norm_rope_flashmla + + swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] + # 未映射槽位(-1, 如 decode 图 warmup 的 HOLD 行: prep 跳过 alloc_swa)对老 triton + # 写入核是显式 no-op;sglang fused 核无负槽位防护(负页偏移=非法访存),mask 到 + # swa HOLD 槽(垃圾桶语义,与 padding 行写入一致)。 + swa_slots = torch.where(swa_slots < 0, torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), swa_slots) + fused_k_norm_rope_flashmla( + kv=kv, + kv_weight=kv_weight, + eps=eps, + freqs_cis=freqs_cis, + positions=positions, + out_loc=swa_slots, + kvcache=self.swa_pool.get_layer_buffer(layer_index), + page_size=self.swa_pool.page_size, + ) + return + def pack_compressed_kv_to_cache(self, layer_index: int, slots: torch.Tensor, comp: torch.Tensor): if comp.shape[0] == 0: return diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 677d3b7dd7..bedf22ee95 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -227,18 +227,68 @@ def apply( ) -> torch.Tensor: raise NotImplementedError("marlin-mxfp4w4a16-b32 is only implemented for fused MoE expert weights") + def _probe_marlin_layout(self, size_n: int, size_k: int, dtype: torch.dtype, device_id: int): + """用零输入走一遍真实的 per-expert repack 路径,探出 marlin 终态布局的形状与类型。 + 只调用 finalize 同款的 vllm 函数,不复刻其内部公式,杜绝形状漂移。结果按维度缓存 + (各 MoE 层同维,全程只探两次: w13 一次、w2 一次)。""" + cache_key = (size_n, size_k, dtype) + cache = getattr(self, "_marlin_layout_cache", None) + if cache is None: + cache = self._marlin_layout_cache = {} + if cache_key in cache: + return cache[cache_key] + + import vllm._custom_ops as ops + from vllm.model_executor.layers.quantization.utils.marlin_utils import ( + get_marlin_input_dtype, + marlin_permute_scales, + ) + from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( + mxfp4_marlin_process_scales, + ) + + input_dtype = get_marlin_input_dtype() + is_a_8bit = input_dtype is not None and input_dtype.itemsize == 1 + device = f"cuda:{device_id}" + qweight = torch.zeros((size_n, size_k // 2), dtype=torch.int8, device=device).view(torch.int32).T.contiguous() + marlin_qweight = ops.gptq_marlin_repack( + b_q_weight=qweight, + perm=torch.empty(0, dtype=torch.int, device=device), + size_k=size_k, + size_n=size_n, + num_bits=4, + is_a_8bit=is_a_8bit, + ) + scale = torch.zeros((size_k // self.block_size, size_n), dtype=dtype, device=device) + marlin_scale = marlin_permute_scales( + s=scale, size_k=size_k, size_n=size_n, group_size=self.block_size, is_a_8bit=is_a_8bit + ) + marlin_scale = mxfp4_marlin_process_scales(marlin_scale, input_dtype=input_dtype) + layout = ( + (tuple(marlin_qweight.shape), marlin_qweight.dtype), + (tuple(marlin_scale.shape), marlin_scale.dtype), + ) + cache[cache_key] = layout + return layout + def _create_weight( self, out_dims: Union[int, List[int]], in_dim: int, dtype: torch.dtype, device_id: int, num_experts: int = 1 ) -> Tuple[WeightPack, List[WeightPack]]: out_dim = sum(out_dims) if isinstance(out_dims, list) else out_dims - assert in_dim % 2 == 0, "MXFP4 packed weight requires even input dimension" assert in_dim % self.block_size == 0, "MXFP4 scale dimension must be divisible by block_size" expert_prefix = (num_experts,) if num_experts > 1 else () + # CPU 暂存区: load_hf_weights 灌入原始预打包 MXFP4,finalize 时 repack 进 CUDA 终态。 weight = torch.empty(expert_prefix + (out_dim, in_dim // 2), dtype=torch.int8, device="cpu") weight_scale = torch.empty( expert_prefix + (out_dim, in_dim // self.block_size), dtype=torch.float8_e8m0fnu, device="cpu" ) mm_param = WeightPack(weight=weight, weight_scale=weight_scale) + # CUDA 终态(marlin 布局)在构造期物化,使 mem manager 的 profile 看到真实权重占用 + # ("构造即分配、load 只灌数"的框架契约,与其它 quant 方法一致;惰性到 finalize 才 + # 进卡会让空卡 profile 把 kv 池撑到挤爆权重加载)。finalize 时 repack 结果拷入。 + (w_shape, w_dtype), (s_shape, s_dtype) = self._probe_marlin_layout(out_dim, in_dim, dtype, device_id) + mm_param.marlin_weight = torch.empty((num_experts,) + w_shape, dtype=w_dtype, device=f"cuda:{device_id}") + mm_param.marlin_weight_scale = torch.empty((num_experts,) + s_shape, dtype=s_dtype, device=f"cuda:{device_id}") mm_param_list = self._split_weight_pack( mm_param, weight_out_dims=out_dims, @@ -267,13 +317,23 @@ class _MXFP4Layer: w13_scale = moe_weight.w13.weight_scale.to(device=device, non_blocking=True).contiguous() w2_scale = moe_weight.w2.weight_scale.to(device=device, non_blocking=True).contiguous() ( - moe_weight.w13.weight, - moe_weight.w2.weight, - moe_weight.w13.weight_scale, - moe_weight.w2.weight_scale, + w13_new, + w2_new, + w13_scale_new, + w2_scale_new, _, _, ) = prepare_moe_mxfp4_layer_for_marlin(layer, w13, w2, w13_scale, w2_scale, None, None) + # repack 结果拷入构造期预分配的 marlin 终态 buffer(与 AWQ marlin 路径同形态), + # CPU 暂存与 repack 临时随引用释放;shape 失配会在 copy_ 处显式报错(探针保证一致)。 + moe_weight.w13.marlin_weight.copy_(w13_new) + moe_weight.w13.marlin_weight_scale.copy_(w13_scale_new) + moe_weight.w2.marlin_weight.copy_(w2_new) + moe_weight.w2.marlin_weight_scale.copy_(w2_scale_new) + moe_weight.w13.weight = moe_weight.w13.marlin_weight + moe_weight.w13.weight_scale = moe_weight.w13.marlin_weight_scale + moe_weight.w2.weight = moe_weight.w2.marlin_weight + moe_weight.w2.weight_scale = moe_weight.w2.marlin_weight_scale def _deepgemm_fp8_nt(a_tuple, b_tuple, out): From c07e38c5979baf3af04ef77773edb7fc90d34f1c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 12 Jun 2026 02:00:32 +0000 Subject: [PATCH 014/214] support fp8 --- .../meta_weights/fused_moe/impl/deepgemm_impl.py | 2 ++ .../meta_weights/fused_moe/impl/marlin_impl.py | 2 ++ .../meta_weights/fused_moe/impl/triton_impl.py | 3 +++ .../triton_kernel/fused_moe/moe_silu_and_mul.py | 11 ++++++++++- .../layer_infer/transformer_layer_infer.py | 4 +++- .../layer_weights/transformer_layer_weight.py | 10 ++++++++-- 6 files changed, 28 insertions(+), 4 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 4d4614c007..72acf2430a 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -76,7 +76,9 @@ def _fused_experts( topk_ids: torch.Tensor, router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, ): + assert clamp_limit is None, "EP deepgemm fused MoE does not support clamp_limit yet" output = fused_experts( hidden_states=input_tensor, w13=w13, diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 0094b09b1c..1fdfd94d0d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -30,7 +30,9 @@ def _fused_experts( topk_ids: torch.Tensor, router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, ): + assert clamp_limit is None, "awq_marlin fused MoE does not support clamp_limit yet" w1_weight, w1_scale, w1_zero_point = w13.weight, w13.weight_scale, w13.weight_zero_point w2_weight, w2_scale, w2_zero_point = w2.weight, w2.weight_scale, w2.weight_zero_point diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index 09ce88e3fd..8967dda34e 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -94,6 +94,7 @@ def _fused_experts( topk_ids: torch.Tensor, router_logits: Optional[torch.Tensor] = None, is_prefill: bool = False, + clamp_limit: Optional[float] = None, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -111,6 +112,7 @@ def _fused_experts( use_fp8_w8a8=use_fp8_w8a8, w1_scale=w13_scale, w2_scale=w2_scale, + limit=clamp_limit, ) return input_tensor @@ -131,6 +133,7 @@ def fused_experts_with_topk( topk_weights=topk_weights, topk_ids=topk_ids, is_prefill=is_prefill, + clamp_limit=clamp_limit, ) def __call__( diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py index 45c7ea73c6..82fc9131c1 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py @@ -24,6 +24,7 @@ def _silu_and_mul_kernel_fast( NEED_MASK: tl.constexpr, layout: tl.constexpr = "blocked", # "blocked" or "interleaved" USE_LIMIT_AND_ALPHA: tl.constexpr = False, + USE_LIMIT_ONLY: tl.constexpr = False, USE_TANH_APPROXIMATE_GELU: tl.constexpr = False, ): stride_input_m = tl.cast(stride_input_m, dtype=tl.int64) @@ -76,6 +77,11 @@ def _silu_and_mul_kernel_fast( mask=mask, ) else: + if USE_LIMIT_ONLY: + # clamped swiglu (DeepSeek-V4 swiglu_limit): clamp 后接标准 silu, + # 无 gpt-oss 的 alpha 缩放与 (up+1)。 + gate = tl.minimum(gate, limit) + up = tl.minimum(tl.maximum(up, -limit), limit) if USE_TANH_APPROXIMATE_GELU: # tanh-approx GELU, matching Gemma's gelu_pytorch_tanh MLP. gate_cubed = gate * gate * gate @@ -124,7 +130,8 @@ def silu_and_mul_fwd( ): assert input.is_contiguous() assert output.is_contiguous() - assert (limit is None and alpha is None) or (limit is not None and alpha is not None) + # limit+alpha: gpt-oss 语义 (up+1)*silu(alpha*gate); 仅 limit: clamp 后标准 silu (DeepSeek-V4) + assert alpha is None or limit is not None stride_input_m = input.stride(0) stride_input_n = input.stride(1) @@ -147,6 +154,7 @@ def silu_and_mul_fwd( while triton.cdiv(size_m, BLOCK_M) > 8192: BLOCK_M *= 2 USE_LIMIT_AND_ALPHA = limit is not None and alpha is not None + USE_LIMIT_ONLY = limit is not None and alpha is None grid = ( triton.cdiv(size_n, BLOCK_N), @@ -171,6 +179,7 @@ def silu_and_mul_fwd( num_warps=num_warps, layout=layout, USE_LIMIT_AND_ALPHA=USE_LIMIT_AND_ALPHA, + USE_LIMIT_ONLY=USE_LIMIT_ONLY, USE_TANH_APPROXIMATE_GELU=ffn_use_tanh_approximate_gelu(), ) return diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 761558c95d..613fb58097 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -595,10 +595,12 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV gw = layer_weight.gate_weight_.mm_param.weight logits = F.linear(x.float(), gw.float()).contiguous() weights, indices = self._select_experts(logits, infer_state, layer_weight) - routed = self._routed_experts(x, weights, indices, layer_weight) + # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts + # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 g = layer_weight.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) u = layer_weight.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) shared = layer_weight.shared_down_.mm((F.silu(g) * u).to(x.dtype)) + routed = self._routed_experts(x, weights, indices, layer_weight) if self.enable_ep_moe: if self.tp_world_size_ > 1: all_reduce( diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index a95299628c..7b2f67e123 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -16,7 +16,8 @@ class DeepseekV4TransformerLayerWeight(TransformerLayerWeight): """Per-layer weights for DeepSeek-V4-Flash. DS4 does not share DS2/DS3.2's ``model.layers.*.self_attn/mlp`` layout. Its attention is - HC + CSA, and routed experts are checkpointed as MXFP4. + HC + CSA, and routed experts are checkpointed as MXFP4 (fp4 release) or + FP8 block-128 (fp8 release, same layout as the dense fp8 weights). """ def __init__(self, layer_num, data_type, network_config, quant_cfg=None): @@ -306,12 +307,17 @@ def _dequant_in_place(self, weights): scale_renames = self._fp8_scale_renames() # Convert every `.scale` belonging to this layer. Weights are loaded incrementally # per safetensors shard, so the paired weight may live in another shard: - # - routed FP4 experts keep `.scale` as-is (matches marlin-mxfp4w4a16-b32's suffix); + # - routed expert `.scale` follows the fused_moe quant method's weight_scale_suffix: + # MXFP4 consumes `.scale` as-is, FP8 DeepGEMM expects `.weight_scale_inv` (rename only); # - FP8 matmul scales only need renaming for DeepGEMM, no weight required; # - FP8 pairs on no-quant paths (wo_a's ROWBMMWeight) are expanded to bf16, # the only case that truly requires weight and scale in the same shard. + expert_scale_suffix = self.experts_.quant_method.weight_scale_suffix for scale_k in [k for k in list(weights.keys()) if k.startswith(p) and k.endswith(".scale")]: if scale_k.startswith(f"{p}ffn.experts."): + if expert_scale_suffix is not None and expert_scale_suffix != "scale": + weights[scale_k[: -len("scale")] + expert_scale_suffix] = weights[scale_k].to(torch.float32) + del weights[scale_k] continue k = scale_k[: -len(".scale")] + ".weight" target = scale_renames.get(k) From ff717061ec5d664218762a785ad752e923f9f50f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 12 Jun 2026 03:23:35 +0000 Subject: [PATCH 015/214] optimize --- .../layer_infer/transformer_layer_infer.py | 115 +++++++++--------- 1 file changed, 57 insertions(+), 58 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 613fb58097..23f742b7d3 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -4,6 +4,7 @@ from lightllm.common.basemodel import TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.distributed.communication_op import all_reduce +from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor @@ -21,27 +22,26 @@ from ..infer_struct import DeepseekV4InferStateInfo -class DeepseekV4TransformerLayerInfer(TransformerLayerInferTpl): +class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): def __init__(self, layer_num, network_config): - super().__init__(layer_num, network_config) - cfg = network_config - self.eps_ = cfg["rms_norm_eps"] - self.hidden = cfg["hidden_size"] - self.n_heads = cfg["num_attention_heads"] - self.head_dim = cfg["head_dim"] - self.rope_dim = cfg["qk_rope_head_dim"] - self.index_n_heads = cfg["index_n_heads"] - self.index_head_dim = cfg["index_head_dim"] - self.index_topk = cfg["index_topk"] - self.o_groups = cfg["o_groups"] - self.o_lora = cfg["o_lora_rank"] - self.hc_mult = cfg["hc_mult"] - self.sinkhorn_iters = cfg["hc_sinkhorn_iters"] - self.hc_eps = cfg["hc_eps"] - self.window = cfg["sliding_window"] - self.compress_ratio = cfg["compress_ratios"][layer_num] - self.is_hash = layer_num < cfg["num_hash_layers"] - self.is_last_layer = layer_num == cfg["n_layer"] - 1 + TransformerLayerInferTpl.__init__(self, layer_num, network_config) + self.eps_ = network_config["rms_norm_eps"] + self.embed_dim_ = network_config["hidden_size"] + self.num_heads = network_config["num_attention_heads"] + self.head_dim_ = network_config["head_dim"] + self.qk_rope_head_dim = network_config["qk_rope_head_dim"] + self.qk_nope_head_dim = self.head_dim_ - self.qk_rope_head_dim + self.v_head_dim = self.head_dim_ + self.index_n_heads = network_config["index_n_heads"] + self.index_head_dim = network_config["index_head_dim"] + self.index_topk = network_config["index_topk"] + self.o_groups = network_config["o_groups"] + self.hc_mult = network_config["hc_mult"] + self.sinkhorn_iters = network_config["hc_sinkhorn_iters"] + self.hc_eps = network_config["hc_eps"] + self.compress_ratio = network_config["compress_ratios"][layer_num] + self.is_hash = layer_num < network_config["num_hash_layers"] + self.is_last_layer = layer_num == network_config["n_layer"] - 1 # complex64 rope table for this layer's variant (sliding / compressed); set by # DeepseekV4TpPartModel._init_to_get_rotary once the tables are built. The full compress # cos/sin tables (compressor entry rope uses entry positions, not token positions) are @@ -49,19 +49,13 @@ def __init__(self, layer_num, network_config): self.freqs_cis = None self.cos_compress_table = None self.sin_compress_table = None - self.topk = cfg["num_experts_per_tok"] - self.route_scale = cfg["routed_scaling_factor"] - self.swiglu_limit = cfg["swiglu_limit"] - self.softmax_scale = self.head_dim ** -0.5 - self.tp_q_heads = self.n_heads // self.tp_world_size_ - self.tp_index_heads = self.index_n_heads // self.tp_world_size_ + self.num_experts_per_tok = network_config["num_experts_per_tok"] + self.routed_scaling_factor = network_config["routed_scaling_factor"] + self.swiglu_limit = network_config["swiglu_limit"] + self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) + self.tp_q_head_num_ = self.num_heads // self.tp_world_size_ + self.tp_index_n_heads = self.index_n_heads // self.tp_world_size_ self.tp_groups = self.o_groups // self.tp_world_size_ - self.tp_q_head_num_ = self.tp_q_heads - self.tp_k_head_num_ = 1 - self.tp_v_head_num_ = 1 - self.tp_o_head_num_ = self.tp_q_heads - self.head_dim_ = self.head_dim - self.embed_dim_ = self.hc_mult * self.hidden self.enable_ep_moe = get_env_start_args().enable_ep_moe self.indexer_score_scale = self.index_head_dim ** -0.5 self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 @@ -72,7 +66,7 @@ def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWe and runs a standalone hc_pre; later layers get (x, residual, post_mix, res_mix) and fuse the previous layer's ffn hc_post with this layer's attn hc_pre.""" if torch.is_tensor(input_embdings): - residual = input_embdings.view(-1, self.hc_mult, self.hidden) + residual = input_embdings.view(-1, self.hc_mult, self.embed_dim_) return hc_pre( residual, layer_weight.hc_attn_fn_.weight, @@ -149,14 +143,19 @@ def _select_rope(self, infer_state: DeepseekV4InferStateInfo): return infer_state.position_cos_compress, infer_state.position_sin_compress return infer_state.position_cos_sliding, infer_state.position_sin_sliding - def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): + def _get_qkv( + self, + input: torch.Tensor, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + ): from sglang.jit_kernel.dsv4 import fused_q_norm_rope - x = self._tpsp_allgather(input=x, infer_state=infer_state) + input = self._tpsp_allgather(input=input, infer_state=infer_state) cos_tok, sin_tok = self._select_rope(infer_state) - T = x.shape[0] - qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(x), eps=self.eps_) - q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_heads, self.head_dim) + T = input.shape[0] + qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(input), eps=self.eps_) + q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, # fused in one sglang dsv4 jit kernel (fp32 norm/rotation, bf16 in between -- same as eager). q = self.alloc_tensor(q_in.shape, dtype=q_in.dtype, device=q_in.device) @@ -167,7 +166,7 @@ def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: Deeps infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, mem_index=infer_state.mem_index, - kv=layer_weight.wkv_.mm(x), + kv=layer_weight.wkv_.mm(input), kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, freqs_cis=self.freqs_cis, @@ -176,7 +175,7 @@ def _get_qkv(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: Deeps return q, qa, cos_tok, sin_tok def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): - # o: [T, tp_q_heads, head_dim] after inverse rope -> grouped low-rank O -> [T, hidden] + # o: [T, tp_q_head_num_, head_dim_] after inverse rope -> grouped low-rank O -> [T, embed_dim_] T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] @@ -185,7 +184,7 @@ def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: Deepsee def _inv_rope(self, o, cos_tok, sin_tok): # in-place; 单张量路径只需要旋转 rope 切片。 - rotary_emb_fwd(o[..., -self.rope_dim :], None, cos_tok, sin_tok, inverse=True) + rotary_emb_fwd(o[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok, inverse=True) return o # ------------------------------------------------------------------ compressor / indexer @@ -196,8 +195,8 @@ def _indexer_q_weight( return None, None cos_tok = infer_state.position_cos_compress sin_tok = infer_state.position_sin_compress - idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_heads, self.index_head_dim) - rotary_emb_fwd(idx_q[..., -self.rope_dim :], None, cos_tok, sin_tok) + idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_n_heads, self.index_head_dim) + rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale return idx_q, idx_weight @@ -231,7 +230,7 @@ def _compressor_weights(self, layer_weight: DeepseekV4TransformerLayerWeight, fo layer_weight.compressor_wgate_.mm_param.weight, layer_weight.compressor_norm_.weight, layer_weight.compressor_ape_.weight, - self.head_dim, + self.head_dim_, ) def _run_compressor_prefill( @@ -286,7 +285,7 @@ def _run_c4_compressor_prefill( wgate, norm, ape, - self.head_dim, + self.head_dim_, self.cos_compress_table, self.sin_compress_table, self.eps_, @@ -337,7 +336,7 @@ def _run_c128_compressor_prefill( norm, ape, self.compress_ratio, - self.head_dim, + self.head_dim_, self.cos_compress_table, self.sin_compress_table, self.eps_, @@ -354,7 +353,7 @@ def _run_c128_compressor_prefill( norm, ape, self.compress_ratio, - self.head_dim, + self.head_dim_, self.cos_compress_table, self.sin_compress_table, self.eps_, @@ -406,7 +405,7 @@ def _run_compressor_decode( wgate, norm, ape, - self.head_dim, + self.head_dim_, self.cos_compress_table, self.sin_compress_table, self.eps_, @@ -424,8 +423,8 @@ def _run_compressor_decode( norm, ape, ratio, - self.head_dim, - self.rope_dim, + self.head_dim_, + self.qk_rope_head_dim, self.cos_compress_table, self.sin_compress_table, self.eps_, @@ -491,7 +490,7 @@ def _context_attention_wrapper_run( infer_state.prefill_cuda_graph_get_current_capture_graph().__enter__() # Same graph-split output handoff as the template, but avoid its dry-run because # DSV4 attention mutates compressor/cache state before returning. - o = self.alloc_tensor((q.shape[0], self.tp_q_heads, self.head_dim), dtype=q.dtype, device=q.device) + o = self.alloc_tensor((q.shape[0], self.tp_q_head_num_, self.head_dim_), dtype=q.dtype, device=q.device) _o = tensor_to_no_ref_tensor(o) def att_func(new_infer_state: DeepseekV4InferStateInfo): @@ -516,7 +515,7 @@ def _context_attention_kernel( "flashmla_kvcache": True, "layer_index": self.layer_num_, "compress_ratio": self.compress_ratio, - "head_dim_v": self.head_dim, + "head_dim_v": self.v_head_dim, "softmax_scale": self.softmax_scale, "q_lora": q_lora, "hidden_states": x, @@ -559,7 +558,7 @@ def _token_attention_kernel( "flashmla_kvcache": True, "layer_index": self.layer_num_, "compress_ratio": self.compress_ratio, - "head_dim_v": self.head_dim, + "head_dim_v": self.v_head_dim, "softmax_scale": self.softmax_scale, "q_lora": q_lora, "hidden_states": x, @@ -588,7 +587,7 @@ def _routed_experts(self, x, weights, indices, layer_weight: DeepseekV4Transform ) def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): - x = x.view(-1, self.hidden) + x = x.view(-1, self.embed_dim_) if not self.enable_ep_moe: x = self._tpsp_allgather(input=x, infer_state=infer_state) @@ -637,16 +636,16 @@ def _select_experts_vllm( else: bias = layer_weight.gate_bias_.weight - weights = self.alloc_tensor((M, self.topk), dtype=torch.float32, device=logits.device) - indices = self.alloc_tensor((M, self.topk), dtype=indices_dtype, device=logits.device) - token_expert_indices = self.alloc_tensor((M, self.topk), dtype=torch.int32, device=logits.device) + weights = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.float32, device=logits.device) + indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=indices_dtype, device=logits.device) + token_expert_indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.int32, device=logits.device) ops.topk_hash_softplus_sqrt( weights, indices, token_expert_indices, logits, True, - self.route_scale, + self.routed_scaling_factor, bias, input_tokens, hash_indices_table, From d7dd6e057acd19053f9ee49b6b983aa401920880 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 12 Jun 2026 03:27:16 +0000 Subject: [PATCH 016/214] fix --- .../layer_infer/transformer_layer_infer.py | 18 +++++++----------- 1 file changed, 7 insertions(+), 11 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 23f742b7d3..cf6c18bfaa 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -152,7 +152,6 @@ def _get_qkv( from sglang.jit_kernel.dsv4 import fused_q_norm_rope input = self._tpsp_allgather(input=input, infer_state=infer_state) - cos_tok, sin_tok = self._select_rope(infer_state) T = input.shape[0] qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(input), eps=self.eps_) q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) @@ -172,21 +171,18 @@ def _get_qkv( freqs_cis=self.freqs_cis, positions=infer_state.position_ids, ) - return q, qa, cos_tok, sin_tok + return q, qa def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): # o: [T, tp_q_head_num_, head_dim_] after inverse rope -> grouped low-rank O -> [T, embed_dim_] + position_cos, position_sin = self._select_rope(infer_state) + rotary_emb_fwd(o[..., -self.qk_rope_head_dim :], None, position_cos, position_sin, inverse=True) T = o.shape[0] o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] o = layer_weight.wo_b_.mm(o) return self._tpsp_reduce(input=o, infer_state=infer_state) - def _inv_rope(self, o, cos_tok, sin_tok): - # in-place; 单张量路径只需要旋转 rope 切片。 - rotary_emb_fwd(o[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok, inverse=True) - return o - # ------------------------------------------------------------------ compressor / indexer def _indexer_q_weight( self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight @@ -468,9 +464,9 @@ def context_attention_forward( # _get_qkv writes the chunk's packed latent into the swa pool (fused kernel) before # attention reads it back via full_to_swa indices (this custom forward bypasses the # tpl _post_cache_kv path). - q, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + q, q_lora = self._get_qkv(x, infer_state, layer_weight) o = self._context_attention_wrapper_run(q, q_lora, x, infer_state, layer_weight) - return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) + return self._get_o(o, infer_state, layer_weight) def _context_attention_wrapper_run( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight @@ -543,9 +539,9 @@ def _context_attention_kernel( def token_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - q, q_lora, cos_tok, sin_tok = self._get_qkv(x, infer_state, layer_weight) + q, q_lora = self._get_qkv(x, infer_state, layer_weight) o = self._token_attention_kernel(q, q_lora, x, infer_state, layer_weight) - return self._get_o(self._inv_rope(o, cos_tok, sin_tok), infer_state, layer_weight) + return self._get_o(o, infer_state, layer_weight) def _token_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight From 3a5dcdc1122d555fe30013759b6e6636beb154f6 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 14 Jun 2026 10:30:50 +0000 Subject: [PATCH 017/214] compress infer --- .../deepseek_v4/layer_infer/compressor.py | 430 ++++++++++++++++++ .../layer_infer/transformer_layer_infer.py | 342 +++----------- .../layer_weights/transformer_layer_weight.py | 15 +- 3 files changed, 503 insertions(+), 284 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 2256ecd1a9..129a93ae0d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -1,4 +1,10 @@ +from dataclasses import dataclass +from typing import Optional + import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice _SGLANG_COMPRESS_ERR = None @@ -7,6 +13,430 @@ _FREQ_CIS_CACHE = {} +@dataclass +class CoreCompressorMetadata: + layer_idx: int + compress_ratio: int + out_slots: torch.Tensor + mem_index: torch.Tensor + state_buffer: torch.Tensor + out_buffer: torch.Tensor + out_page_size: int + position_ids: torch.Tensor + b_req_idx: torch.Tensor + b_seq_len: torch.Tensor + b_ready_cache_len: Optional[torch.Tensor] + b_q_start_loc: Optional[torch.Tensor] + req_to_token_indexs: torch.Tensor + full_to_swa_indexs: torch.Tensor + token_to_batch_idx: Optional[torch.Tensor] + kv_score: Optional[torch.Tensor] + is_prefill: bool + + +@triton.jit +def _add_ape_to_kv_score_kernel( + kv_score, + kv_score_stride0, + kv_score_stride1, + ape, + ape_stride0, + positions, + STATE_WIDTH: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + BLOCK: tl.constexpr, +): + token_idx = tl.program_id(0) + offs = tl.arange(0, BLOCK) + mask = offs < STATE_WIDTH + + position = tl.load(positions + token_idx) + ape_row = position % COMPRESS_RATIO + score = tl.load(kv_score + token_idx * kv_score_stride0 + (STATE_WIDTH + offs) * kv_score_stride1, mask=mask) + ape_value = tl.load(ape + ape_row * ape_stride0 + offs, mask=mask) + tl.store( + kv_score + token_idx * kv_score_stride0 + (STATE_WIDTH + offs) * kv_score_stride1, + score + ape_value, + mask=mask, + ) + return + + +@triton.jit +def _save_partial_states_kernel( + kv_score, + kv_score_stride0, + kv_score_stride1, + positions, + token_to_batch_idx, + b_req_idx, + b_seq_len, + mem_index, + full_to_swa, + state_buffer, + STATE_WIDTH: tl.constexpr, + STATE_LAST_DIM: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + IS_C4: tl.constexpr, + IS_PREFILL: tl.constexpr, + SWA_PAGE_SIZE: tl.constexpr, + C4_STATE_RING: tl.constexpr, + BLOCK: tl.constexpr, +): + token_idx = tl.program_id(0) + batch_idx = tl.load(token_to_batch_idx + token_idx) if IS_PREFILL else token_idx + position = tl.load(positions + token_idx) + seq_len = tl.load(b_seq_len + batch_idx) + + if IS_C4: + same_page_next = (position % SWA_PAGE_SIZE) + C4_STATE_RING < SWA_PAGE_SIZE + if same_page_next and position + C4_STATE_RING < seq_len: + return + full_slot = tl.load(mem_index + token_idx).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + if swa_slot < 0: + return + state_row = (swa_slot // SWA_PAGE_SIZE) * C4_STATE_RING + (swa_slot % C4_STATE_RING) + else: + if position + COMPRESS_RATIO < seq_len: + return + req_idx = tl.load(b_req_idx + batch_idx).to(tl.int64) + state_row = req_idx * COMPRESS_RATIO + (position % COMPRESS_RATIO) + + offs = tl.arange(0, BLOCK) + mask = offs < STATE_WIDTH + kv = tl.load(kv_score + token_idx * kv_score_stride0 + offs * kv_score_stride1, mask=mask) + score = tl.load(kv_score + token_idx * kv_score_stride0 + (STATE_WIDTH + offs) * kv_score_stride1, mask=mask) + state_base = state_buffer + state_row * STATE_LAST_DIM + tl.store(state_base + offs, kv, mask=mask) + tl.store(state_base + STATE_WIDTH + offs, score, mask=mask) + return + + +@triton.jit +def _fused_compress_norm_rope_insert_kernel( + kv_score, + kv_score_stride0, + kv_score_stride1, + state_buffer, + positions, + token_to_batch_idx, + b_req_idx, + b_seq_len, + b_ready_cache_len, + b_q_start_loc, + req_to_token, + req_to_token_stride0, + full_to_swa, + out_slots, + norm_weight, + rms_eps, + cos_table, + cos_stride0, + sin_table, + sin_stride0, + out_buffer, + HEAD_DIM: tl.constexpr, + STATE_WIDTH: tl.constexpr, + STATE_LAST_DIM: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, + WINDOW_SIZE: tl.constexpr, + IS_C4: tl.constexpr, + IS_PREFILL: tl.constexpr, + SWA_PAGE_SIZE: tl.constexpr, + C4_STATE_RING: tl.constexpr, + ROPE_HEAD_DIM: tl.constexpr, + FP8_MAX: tl.constexpr, + SCALE_MIN: tl.constexpr, + NOPE_DIM: tl.constexpr, + QUANT_BLOCK: tl.constexpr, + SCALE_BYTES: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BYTES_PER_PAGE: tl.constexpr, + BLOCK: tl.constexpr, +): + token_idx = tl.program_id(0) + out_slot = tl.load(out_slots + token_idx).to(tl.int64) + if out_slot < 0: + return + + position = tl.load(positions + token_idx) + if (position + 1) % COMPRESS_RATIO != 0: + return + + batch_idx = tl.load(token_to_batch_idx + token_idx) if IS_PREFILL else token_idx + req_idx = tl.load(b_req_idx + batch_idx).to(tl.int64) + seq_len = tl.load(b_seq_len + batch_idx) + if IS_PREFILL: + ready_len = tl.load(b_ready_cache_len + batch_idx) + q_start = tl.load(b_q_start_loc + batch_idx) + else: + ready_len = position + q_start = token_idx + + token_offsets = tl.arange(0, WINDOW_SIZE) + start = position - WINDOW_SIZE + 1 + gather_pos = start + token_offsets + valid_pos = (gather_pos >= 0) & (gather_pos < seq_len) + use_current = (gather_pos >= ready_len) & valid_pos if IS_PREFILL else gather_pos == position + current_idx = q_start + (gather_pos - ready_len) if IS_PREFILL else token_idx + token_offsets * 0 + + if IS_C4: + full_slot = tl.load( + req_to_token + req_idx * req_to_token_stride0 + gather_pos, + mask=valid_pos & (~use_current), + other=0, + ).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot, mask=valid_pos & (~use_current), other=-1).to(tl.int64) + state_row = (swa_slot // SWA_PAGE_SIZE) * C4_STATE_RING + (swa_slot % C4_STATE_RING) + state_valid = valid_pos & (~use_current) & (swa_slot >= 0) + head_offset = tl.where(token_offsets >= COMPRESS_RATIO, HEAD_DIM, 0) + else: + state_row = req_idx * COMPRESS_RATIO + (gather_pos % COMPRESS_RATIO) + state_valid = valid_pos & (~use_current) + head_offset = token_offsets * 0 + + offs = tl.arange(0, BLOCK) + dim_mask = offs < HEAD_DIM + current_mask = use_current[:, None] & dim_mask[None, :] + state_mask = state_valid[:, None] & dim_mask[None, :] + + cur_kv = tl.load( + kv_score + current_idx[:, None] * kv_score_stride0 + (head_offset[:, None] + offs[None, :]) * kv_score_stride1, + mask=current_mask, + other=0.0, + ) + cur_score = tl.load( + kv_score + + current_idx[:, None] * kv_score_stride0 + + (STATE_WIDTH + head_offset[:, None] + offs[None, :]) * kv_score_stride1, + mask=current_mask, + other=float("-inf"), + ) + state_kv = tl.load( + state_buffer + state_row[:, None] * STATE_LAST_DIM + head_offset[:, None] + offs[None, :], + mask=state_mask, + other=0.0, + ) + state_score = tl.load( + state_buffer + state_row[:, None] * STATE_LAST_DIM + STATE_WIDTH + head_offset[:, None] + offs[None, :], + mask=state_mask, + other=float("-inf"), + ) + + kv = tl.where(current_mask, cur_kv, state_kv) + score = tl.where(current_mask, cur_score, state_score) + score = tl.softmax(score, dim=0) + compressed_kv = tl.sum(kv * score, axis=0) + + rms_w = tl.load(norm_weight + offs, mask=dim_mask, other=0.0) + variance = tl.sum(compressed_kv * compressed_kv, axis=0) / HEAD_DIM + rrms = tl.rsqrt(variance + rms_eps) + normed = compressed_kv * rrms * rms_w + + num_pairs: tl.constexpr = BLOCK // 2 + nope_pairs: tl.constexpr = NOPE_DIM // 2 + pair_2d = tl.reshape(normed, (num_pairs, 2)) + even, odd = tl.split(pair_2d) + pair_idx = tl.arange(0, num_pairs) + rope_pair_local = pair_idx - nope_pairs + is_rope_pair = rope_pair_local >= 0 + cs_idx = tl.maximum(rope_pair_local, 0) + compressed_pos = (position // COMPRESS_RATIO) * COMPRESS_RATIO + cos_v = tl.load(cos_table + compressed_pos * cos_stride0 + cs_idx, mask=is_rope_pair, other=1.0) + sin_v = tl.load(sin_table + compressed_pos * sin_stride0 + cs_idx, mask=is_rope_pair, other=0.0) + new_even = even * cos_v - odd * sin_v + new_odd = odd * cos_v + even * sin_v + rotated = tl.interleave(new_even, new_odd) + + page = out_slot // PAGE_SIZE + token_in_page = out_slot % PAGE_SIZE + data_base = page * BYTES_PER_PAGE + token_in_page * (NOPE_DIM + ROPE_HEAD_DIM * 2) + scale_base = page * BYTES_PER_PAGE + PAGE_SIZE * (NOPE_DIM + ROPE_HEAD_DIM * 2) + token_in_page * SCALE_BYTES + + n_quant_blocks: tl.constexpr = BLOCK // QUANT_BLOCK + n_nope_blocks: tl.constexpr = NOPE_DIM // QUANT_BLOCK + quant_input = normed.to(tl.bfloat16).to(tl.float32) + quant_2d = tl.reshape(quant_input, (n_quant_blocks, QUANT_BLOCK)) + abs_2d = tl.abs(quant_2d) + block_absmax = tl.max(abs_2d, axis=1) + scale_exp = tl.ceil(libdevice.log2(tl.maximum(block_absmax / FP8_MAX, SCALE_MIN))).to(tl.int32) + scale = ((scale_exp + 127) << 23).to(tl.float32, bitcast=True) + kv_fp8 = tl.clamp(quant_2d / scale[:, None], -FP8_MAX, FP8_MAX).to(tl.float8e4nv) + kv_u8 = tl.reshape(kv_fp8.to(tl.uint8, bitcast=True), (BLOCK,)) + tl.store(out_buffer + data_base + offs, kv_u8, mask=offs < NOPE_DIM) + + scale_idx = tl.arange(0, SCALE_BYTES) + scale_bytes = tl.where(scale_idx < n_nope_blocks, scale_exp + 127, 0).to(tl.uint8) + tl.store(out_buffer + scale_base + scale_idx, scale_bytes) + + rope_local = offs - NOPE_DIM + rope_mask = (offs >= NOPE_DIM) & dim_mask + rope_ptr = (out_buffer + data_base + NOPE_DIM).to(tl.pointer_type(tl.bfloat16)) + tl.store(rope_ptr + rope_local, rotated.to(tl.bfloat16), mask=rope_mask) + return + + +def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int): + if compress_ratio == 0: + return None + + mem_manager = infer_state.mem_manager + if compress_ratio == 4: + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + state_buffer = mem_manager.get_c4_state_buffer(layer_idx) + out_pool = mem_manager.c4_pool + elif compress_ratio == 128: + out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] + state_buffer = infer_state.req_manager.get_compress_state_pool(layer_idx) + out_pool = mem_manager.c128_pool + else: + raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") + + token_to_batch_idx = infer_state.b_req_idx + if infer_state.is_prefill: + token_to_batch_idx = getattr(infer_state, "_dsv4_token_to_batch_idx", None) + if token_to_batch_idx is None or token_to_batch_idx.numel() != infer_state.position_ids.numel(): + q_lens = (infer_state.b_seq_len - infer_state.b_ready_cache_len).to(torch.long) + batch_idx = torch.arange(infer_state.b_req_idx.shape[0], device=infer_state.b_req_idx.device) + token_to_batch_idx = torch.repeat_interleave(batch_idx, q_lens).to(torch.int32) + infer_state._dsv4_token_to_batch_idx = token_to_batch_idx + + return CoreCompressorMetadata( + layer_idx=layer_idx, + compress_ratio=compress_ratio, + out_slots=out_slots, + mem_index=infer_state.mem_index, + state_buffer=state_buffer, + out_buffer=mem_manager.get_compressed_kv_buffer(layer_idx), + out_page_size=out_pool.page_size, + position_ids=infer_state.position_ids, + b_req_idx=infer_state.b_req_idx, + b_seq_len=infer_state.b_seq_len, + b_ready_cache_len=infer_state.b_ready_cache_len, + b_q_start_loc=infer_state.b_q_start_loc, + req_to_token_indexs=infer_state.req_manager.req_to_token_indexs, + full_to_swa_indexs=mem_manager.full_to_swa_indexs, + token_to_batch_idx=token_to_batch_idx, + kv_score=None, + is_prefill=infer_state.is_prefill, + ) + + +def prepare_partial_states( + *, + kv_score: torch.Tensor, + metadata: Optional[CoreCompressorMetadata], + ape: torch.Tensor, + compress_ratio: int, +): + if metadata is None or kv_score.shape[0] == 0: + return + state_width = kv_score.shape[-1] // 2 + _add_ape_to_kv_score_kernel[(kv_score.shape[0],)]( + kv_score, + kv_score.stride(0), + kv_score.stride(1), + ape, + ape.stride(0), + metadata.position_ids, + STATE_WIDTH=state_width, + COMPRESS_RATIO=compress_ratio, + BLOCK=triton.next_power_of_2(state_width), + num_warps=4, + ) + return + + +def fused_compress( + *, + kv_score: torch.Tensor, + metadata: Optional[CoreCompressorMetadata], + norm_weight: torch.Tensor, + ape: torch.Tensor, + eps: float, + head_dim: int, + qk_rope_head_dim: int, + compress_ratio: int, + cos_table: torch.Tensor, + sin_table: torch.Tensor, +): + if metadata is None or kv_score.shape[0] == 0: + return + + state_width = kv_score.shape[-1] // 2 + state_last_dim = metadata.state_buffer.shape[-1] + is_c4 = compress_ratio == 4 + block_state = triton.next_power_of_2(state_width) + block_head = triton.next_power_of_2(head_dim) + + _fused_compress_norm_rope_insert_kernel[(kv_score.shape[0],)]( + kv_score, + kv_score.stride(0), + kv_score.stride(1), + metadata.state_buffer, + metadata.position_ids, + metadata.token_to_batch_idx, + metadata.b_req_idx, + metadata.b_seq_len, + metadata.b_ready_cache_len if metadata.b_ready_cache_len is not None else metadata.b_seq_len, + metadata.b_q_start_loc if metadata.b_q_start_loc is not None else metadata.b_seq_len, + metadata.req_to_token_indexs, + metadata.req_to_token_indexs.stride(0), + metadata.full_to_swa_indexs, + metadata.out_slots, + norm_weight, + eps, + cos_table, + cos_table.stride(0), + sin_table, + sin_table.stride(0), + metadata.out_buffer, + HEAD_DIM=head_dim, + STATE_WIDTH=state_width, + STATE_LAST_DIM=state_last_dim, + COMPRESS_RATIO=compress_ratio, + WINDOW_SIZE=compress_ratio * (2 if is_c4 else 1), + IS_C4=is_c4, + IS_PREFILL=metadata.is_prefill, + SWA_PAGE_SIZE=128, + C4_STATE_RING=8, + ROPE_HEAD_DIM=qk_rope_head_dim, + FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, + SCALE_MIN=1e-4, + NOPE_DIM=head_dim - qk_rope_head_dim, + QUANT_BLOCK=64, + SCALE_BYTES=(head_dim - qk_rope_head_dim) // 64 + 1, + PAGE_SIZE=metadata.out_page_size, + BYTES_PER_PAGE=metadata.out_buffer.shape[-1], + BLOCK=block_head, + num_warps=4, + ) + + _save_partial_states_kernel[(kv_score.shape[0],)]( + kv_score, + kv_score.stride(0), + kv_score.stride(1), + metadata.position_ids, + metadata.token_to_batch_idx, + metadata.b_req_idx, + metadata.b_seq_len, + metadata.mem_index, + metadata.full_to_swa_indexs, + metadata.state_buffer, + STATE_WIDTH=state_width, + STATE_LAST_DIM=state_last_dim, + COMPRESS_RATIO=compress_ratio, + IS_C4=is_c4, + IS_PREFILL=metadata.is_prefill, + SWA_PAGE_SIZE=128, + C4_STATE_RING=8, + BLOCK=block_state, + num_warps=4, + ) + return + + def _load_sglang_compressor(): global _SGLANG_COMPRESS_ERR, _SGLANG_COMPRESS_MOD, _SGLANG_LINEAR_BF16_FP32 if _SGLANG_COMPRESS_MOD is not None: diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index cf6c18bfaa..1b5c5f4c3f 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -9,15 +9,9 @@ from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post -from .compressor import ( - compressor_prefill_state, - compressor_decode_step_single, - compressor_decode_step_batch, - compressor_paged_prefill, - compressor_paged_decode_batch, - paged_prefill_compress_data, - paged_decode_state_slots, -) +from .compressor import fused_compress as fused_compress_op +from .compressor import prepare_partial_states +from .compressor import prepare_compress_states from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo @@ -59,6 +53,9 @@ def __init__(self, layer_num, network_config): self.enable_ep_moe = get_env_start_args().enable_ep_moe self.indexer_score_scale = self.index_head_dim ** -0.5 self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 + self.compressor = CompressorInfer( + layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ + ) # ------------------------------------------------------------------ forward (HC-threaded) def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): @@ -196,267 +193,6 @@ def _indexer_q_weight( idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale return idx_q, idx_weight - def _gather_compress_slots(self, infer_state: DeepseekV4InferStateInfo, req, entry_start, entry_count): - """组末 token 的 full 槽位 -> 压缩槽(条目 [entry_start, entry_start+entry_count))。 - 槽位已由 prep 阶段(prepare_*_compress_slots)分配并 scatter 进 full_to_c4/c128_indexs。""" - ratio = self.compress_ratio - mem = infer_state.mem_manager - mapping = mem.full_to_c4_indexs if ratio == 4 else mem.full_to_c128_indexs - last = entry_start + entry_count - ends = infer_state.req_manager.req_to_token_indexs[req, ratio - 1 : last * ratio : ratio][entry_start:] - return mapping[ends.long()] - - def _write_compressed_kv(self, infer_state: DeepseekV4InferStateInfo, req, entry_start, comp): - slots = self._gather_compress_slots(infer_state, req, entry_start, comp.shape[0]) - if comp.shape[0]: - infer_state.mem_manager.pack_compressed_kv_to_cache(self.layer_num_, slots, comp) - return slots - - def _compressor_weights(self, layer_weight: DeepseekV4TransformerLayerWeight, for_indexer: bool): - if for_indexer: - return ( - layer_weight.idx_cmp_wkv_.mm_param.weight, - layer_weight.idx_cmp_wgate_.mm_param.weight, - layer_weight.idx_cmp_norm_.weight, - layer_weight.idx_cmp_ape_.weight, - self.index_head_dim, - ) - return ( - layer_weight.compressor_wkv_.mm_param.weight, - layer_weight.compressor_wgate_.mm_param.weight, - layer_weight.compressor_norm_.weight, - layer_weight.compressor_ape_.weight, - self.head_dim_, - ) - - def _run_compressor_prefill( - self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - """Per-request compressor for the prefill chunk. Runs as part of the deferred attention - func, before the attention metadata gathers the slot mappings. - - c4: paged state (swa-page-derived group slots, translation #3) — one fused extend-aware - call per request; the (write_loc, extra_data, plan) tuple is layer-independent and cached - on infer_state across all c4 layers. c128: req-keyed state (zero at every 128 boundary by - construction, nothing cache-resident), original jit paths.""" - if not self.compress_ratio: - return - if self.compress_ratio == 4: - self._run_c4_compressor_prefill(x, infer_state, layer_weight) - else: - self._run_c128_compressor_prefill(x, infer_state, layer_weight) - return - - def _run_c4_compressor_prefill( - self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - rm = infer_state.req_manager - mem = infer_state.mem_manager - wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) - iwkv, iwgate, inorm, iape, _ = self._compressor_weights(layer_weight, for_indexer=True) - state_buf = mem.get_c4_state_buffer(self.layer_num_) - idx_state_buf = mem.get_c4_indexer_state_buffer(self.layer_num_) - data_cache = getattr(infer_state, "_dsv4_c4_prefill_data", None) - if data_cache is None: - data_cache = {} - infer_state._dsv4_c4_prefill_data = data_cache - b_req = infer_state.b_req_idx.tolist() - starts = infer_state.b_q_start_loc.tolist() - lens = infer_state.b_q_seq_len.tolist() - ready_lens = infer_state.b_ready_cache_len.tolist() - for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): - if req == rm.HOLD_REQUEST_ID or ln == 0: - continue - seq_len = ready_len + ln - data = data_cache.get(req) - if data is None: - data = paged_prefill_compress_data( - rm.req_to_token_indexs, mem.full_to_swa_indexs, req, ready_len, seq_len, ring=8 - ) - data_cache[req] = data - x_r = x[st : st + ln] - comp = compressor_paged_prefill( - x_r, - wkv, - wgate, - norm, - ape, - self.head_dim_, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - state_buf, - data, - ready_len, - seq_len, - ) - slots = self._write_compressed_kv(infer_state, req, ready_len // 4, comp) - idx_comp = compressor_paged_prefill( - x_r, - iwkv, - iwgate, - inorm, - iape, - self.index_head_dim, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - idx_state_buf, - data, - ready_len, - seq_len, - ) - if idx_comp.shape[0]: - infer_state.mem_manager.pack_indexer_k_to_cache(self.layer_num_, slots, idx_comp) - return - - def _run_c128_compressor_prefill( - self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - rm = infer_state.req_manager - wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) - b_req = infer_state.b_req_idx.tolist() - starts = infer_state.b_q_start_loc.tolist() - lens = infer_state.b_q_seq_len.tolist() - ready_lens = infer_state.b_ready_cache_len.tolist() - for req, st, ln, ready_len in zip(b_req, starts, lens, ready_lens): - if req == rm.HOLD_REQUEST_ID: - continue - x_r = x[st : st + ln] - state_pool = rm.get_compress_state_pool_for_req(self.layer_num_, req) - if ready_len == 0: - comp = compressor_prefill_state( - x_r, - wkv, - wgate, - norm, - ape, - self.compress_ratio, - self.head_dim_, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - state_pool, - ) - self._write_compressed_kv(infer_state, req, 0, comp) - else: - for j in range(ln): - start_pos = ready_len + j - entry = compressor_decode_step_single( - x_r[j], - wkv, - wgate, - norm, - ape, - self.compress_ratio, - self.head_dim_, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - state_pool, - start_pos, - ) - if entry is not None: - entry_start = (start_pos + 1) // self.compress_ratio - 1 - self._write_compressed_kv(infer_state, req, entry_start, entry.unsqueeze(0)) - return - - def _run_compressor_decode( - self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - """Batched decode compressor (cuda-graph safe): state update for every request, cache write - masked to the pool HOLD slot unless this token completes a window. Compressed-cache slots - were pre-allocated by prepare_decode_compress_slots in the prep phase. - - c4: paged state — group slots derived from full_to_swa (translation #3) via pure tensor - ops (graph-safe), shared across all c4 layers per step. c128: req-keyed state.""" - if not self.compress_ratio: - return - rm = infer_state.req_manager - mem = infer_state.mem_manager - req = infer_state.b_req_idx - ratio = self.compress_ratio - wkv, wgate, norm, ape, _ = self._compressor_weights(layer_weight, for_indexer=False) - - if ratio == 4: - mapping, hold = mem.full_to_c4_indexs, mem.c4_pool.HOLD_TOKEN_MEMINDEX - slot_meta = getattr(infer_state, "_dsv4_c4_decode_slots", None) - if slot_meta is None: - slot_meta = paged_decode_state_slots( - rm.req_to_token_indexs, - mem.full_to_swa_indexs, - req, - infer_state.b_seq_len, - page_size=128, - ring=8, - ratio=4, - hold_req_id=rm.HOLD_REQUEST_ID, - num_swa_pages=mem.swa_num_pages, - ) - infer_state._dsv4_c4_decode_slots = slot_meta - write_slot, overlap_slot = slot_meta - entry, should = compressor_paged_decode_batch( - x, - wkv, - wgate, - norm, - ape, - self.head_dim_, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - mem.get_c4_state_buffer(self.layer_num_), - write_slot, - overlap_slot, - infer_state.b_seq_len, - ) - else: - mapping, hold = mem.full_to_c128_indexs, mem.c128_pool.HOLD_TOKEN_MEMINDEX - entry, should = compressor_decode_step_batch( - x, - wkv, - wgate, - norm, - ape, - ratio, - self.head_dim_, - self.qk_rope_head_dim, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - rm.get_compress_state_pool(self.layer_num_), - req, - infer_state.b_seq_len.long() - 1, - ) - - should = should & (req != rm.HOLD_REQUEST_ID) - # 本步 token 即组末 token(should 为真时),其 full 槽 = mem_index,映射在 prep 已 scatter。 - slots = mapping[infer_state.mem_index.long()].long() - slots = torch.where(should, slots, torch.full_like(slots, hold)) - mem.pack_compressed_kv_to_cache(self.layer_num_, slots, entry) - - if ratio == 4: - iwkv, iwgate, inorm, iape, _ = self._compressor_weights(layer_weight, for_indexer=True) - idx_entry, idx_should = compressor_paged_decode_batch( - x, - iwkv, - iwgate, - inorm, - iape, - self.index_head_dim, - self.cos_compress_table, - self.sin_compress_table, - self.eps_, - mem.get_c4_indexer_state_buffer(self.layer_num_), - write_slot, - overlap_slot, - infer_state.b_seq_len, - ) - idx_should = idx_should & (req != rm.HOLD_REQUEST_ID) - idx_slots = torch.where(idx_should, slots, torch.full_like(slots, hold)) - mem.pack_indexer_k_to_cache(self.layer_num_, idx_slots, idx_entry) - return - # ------------------------------------------------------------------ attention (prefill) def context_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight @@ -503,7 +239,8 @@ def att_func(new_infer_state: DeepseekV4InferStateInfo): def _context_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - self._run_compressor_prefill(x, infer_state, layer_weight) + self.compressor.prepare_states(x, infer_state, layer_weight) + self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) att_control = AttControl( nsa_prefill=True, @@ -546,7 +283,8 @@ def token_attention_forward( def _token_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - self._run_compressor_decode(x, infer_state, layer_weight) + self.compressor.prepare_states(x, infer_state, layer_weight) + self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) att_control = AttControl( nsa_decode=True, @@ -647,3 +385,63 @@ def _select_experts_vllm( hash_indices_table, ) return weights, indices.long() + + +class CompressorInfer: + def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): + super().__init__() + self.layer_idx_ = layer_idx + self.network_config_ = network_config + self.tp_world_size_ = tp_world_size + self.compress_ratio = network_config["compress_ratios"][layer_idx] + self.head_dim = network_config["head_dim"] + self.index_head_dim = network_config["index_head_dim"] + self.qk_rope_head_dim = network_config["qk_rope_head_dim"] + self.eps = network_config["rms_norm_eps"] + self._metadata = None + + def prepare_states( + self, + x: torch.Tensor, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + ): + self._metadata = prepare_compress_states( + infer_state=infer_state, + layer_idx=self.layer_idx_, + compress_ratio=self.compress_ratio, + ) + if self._metadata is not None: + self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm(x).float() + prepare_partial_states( + kv_score=self._metadata.kv_score, + metadata=self._metadata, + ape=layer_weight.compressor_ape_.weight, + compress_ratio=self.compress_ratio, + ) + return self._metadata + + def fused_compress( + self, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + cos_table: torch.Tensor, + sin_table: torch.Tensor, + ): + if self.compress_ratio == 0: + return None + metadata = self._metadata + if metadata is None: + raise RuntimeError("DeepSeek-V4 compressor.prepare_states must run before fused_compress") + return fused_compress_op( + kv_score=metadata.kv_score, + metadata=metadata, + norm_weight=layer_weight.compressor_norm_.weight, + ape=layer_weight.compressor_ape_.weight, + eps=self.eps, + head_dim=self.head_dim, + qk_rope_head_dim=self.qk_rope_head_dim, + compress_ratio=self.compress_ratio, + cos_table=cos_table, + sin_table=sin_table, + ) diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 7b2f67e123..d42b20f6e5 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -121,19 +121,10 @@ def _init_compressor(self): coff = 2 if ratio == 4 else 1 # wkv/wgate are bf16 (no scale) and replicated (single KV head). - self.compressor_wkv_ = ROWMMWeight( + self.compressor_wkv_gate_ = ROWMMWeight( in_dim=self.hidden, - out_dims=[coff * head_dim], - weight_names=f"{prefix}.wkv.weight", - data_type=self.data_type_, - quant_method=None, - tp_rank=0, - tp_world_size=1, - ) - self.compressor_wgate_ = ROWMMWeight( - in_dim=self.hidden, - out_dims=[coff * head_dim], - weight_names=f"{prefix}.wgate.weight", + out_dims=[coff * head_dim, coff * head_dim], + weight_names=[f"{prefix}.wkv.weight", f"{prefix}.wgate.weight"], data_type=self.data_type_, quant_method=None, tp_rank=0, From d76450f9c92af47857f33f3615cfdfbcc48ef55b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 14 Jun 2026 13:27:05 +0000 Subject: [PATCH 018/214] add c128 to mem_manager --- .../deepseek4_mem_manager.py | 60 +++++++--- lightllm/common/req_manager.py | 42 +------ .../deepseek_v4/layer_infer/compressor.py | 104 +++++++++++------- 3 files changed, 116 insertions(+), 90 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index f87a11704d..8d172ec758 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -31,9 +31,9 @@ DSV4_SWA_PAGE_SIZE = 128 DSV4_C4_PAGE_SIZE = 64 DSV4_C128_PAGE_SIZE = 2 -# c4 compressor state ring(overlap 对: 每页 2 个分组槽 × ratio 4 行)。c128 state 在 128 边界 -# 自然归零(在线聚合),无缓存常驻需求,保持 req 键控,不进 swa 派生池。 +# compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 离线聚合为每页 1 组。 DSV4_C4_STATE_RING = 8 +DSV4_C128_STATE_RING = 128 # swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 同值)。 # v5 的 swa 压力阀(借页/驱逐)已覆盖 radix 树与准入波次的瞬时增长,结构性预算 # (max_req×window + batch_max_tokens 余量)另行叠加,0.1 仅作 full 池比例下限。 @@ -154,6 +154,7 @@ def __init__( max_request_num: Optional[int] = None, sliding_window: Optional[int] = None, swa_extra_token_num: int = 0, + swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, ): @@ -174,6 +175,7 @@ def __init__( # 活跃窗口(max_request_num * sliding_window)之外的余量: 在途 prefill chunk 的瞬时占用 # (出窗槽位要到下一次 prep 才回收) + radix cache 持有的窗口尾部。 self.swa_extra_token_num = int(swa_extra_token_num) + self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) self.layer_to_c4_idx = {} @@ -200,7 +202,7 @@ def _planned_swa_size(self, full_size: int) -> int: if self.max_request_num is None or self.sliding_window is None: return _ceil_div(full_size, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE cap = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num - cap = max(cap, int(full_size * DSV4_SWA_FULL_TOKENS_RATIO)) + cap = max(cap, int(full_size * self.swa_full_tokens_ratio)) cap = max(1, min(full_size, cap)) return _ceil_div(cap, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE @@ -209,18 +211,32 @@ def _slab_bytes_per_slot(page_size: int, data_bytes: int, scale_bytes: int, alig bytes_per_page = _ceil_div(page_size * (data_bytes + scale_bytes), align_bytes) * align_bytes return bytes_per_page / page_size - def _c4_state_bytes_per_swa_slot(self) -> float: - """c4 compressor state(attention + indexer,swa 页派生寻址)摊到每个 swa 槽的字节数。""" - if self.n_c4 == 0: - return 0.0 - per_page = DSV4_C4_STATE_RING * (4 * self.head_dim + 4 * self.indexer_head_dim) * 4 # fp32 - return per_page * self.n_c4 / DSV4_SWA_PAGE_SIZE + @staticmethod + def _paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: + rows = num_swa_pages * ring + ring + 1 + return _ceil_div(rows, ratio) * ratio + + @staticmethod + def _init_state_sentinel(buffer: torch.Tensor) -> None: + half = buffer.shape[-1] // 2 + buffer[:, -1, :half].zero_() + buffer[:, -1, half:].fill_(float("-inf")) + return + + def _paged_state_bytes_per_swa_slot(self) -> float: + """c4/c128 compressor state(swa 页派生寻址)摊到每个 swa 槽的字节数。""" + per_page = 0.0 + if self.n_c4 > 0: + per_page += DSV4_C4_STATE_RING * (4 * self.head_dim + 4 * self.indexer_head_dim) * 4 * self.n_c4 + if self.n_c128 > 0: + per_page += DSV4_C128_STATE_RING * (2 * self.head_dim) * 4 * self.n_c128 + return per_page / DSV4_SWA_PAGE_SIZE def _swa_slot_bytes(self) -> float: per_layer = self._slab_bytes_per_slot( DSV4_SWA_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES ) - return per_layer * self.layer_num + self._c4_state_bytes_per_swa_slot() + return per_layer * self.layer_num + self._paged_state_bytes_per_swa_slot() def _compressed_cell_size(self) -> float: """每个 full token 摊到压缩池上的精确字节数(按 page-slab 对齐后)。""" @@ -259,10 +275,10 @@ def profile_size(self, mem_fraction): self.size = max(1, int(available_bytes / full_cell)) else: size_budget = max(1, int((available_bytes - swa_slot_bytes * swa_budget) / compressed_cell)) - if size_budget * DSV4_SWA_FULL_TOKENS_RATIO > swa_budget: + if size_budget * self.swa_full_tokens_ratio > swa_budget: # 比例下限生效(_planned_swa_size 会取 ratio*full),按该机制反解 full。 self.size = max( - 1, int(available_bytes / (swa_slot_bytes * DSV4_SWA_FULL_TOKENS_RATIO + compressed_cell)) + 1, int(available_bytes / (swa_slot_bytes * self.swa_full_tokens_ratio + compressed_cell)) ) else: self.size = size_budget @@ -321,6 +337,9 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.c4_allocator: Optional[KvCacheAllocator] = None self.c128_pool: Optional[PackedPagePool] = None self.c128_allocator: Optional[KvCacheAllocator] = None + self.c4_state_buffer: Optional[torch.Tensor] = None + self.c4_indexer_state_buffer: Optional[torch.Tensor] = None + self.c128_state_buffer: Optional[torch.Tensor] = None # 压缩槽映射: 键 = 组末 token(位置 (g+1)%ratio==0)的 full 槽位,值 = 压缩池槽位。 # 与 full_to_swa_indexs 同构: radix 持有 full 槽 => 映射行存活,free 级联回收。 self.full_to_c4_indexs: Optional[torch.Tensor] = None @@ -350,8 +369,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): # 生灭 -> radix 命中零拷贝续算。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), # 取整到 ratio;末行哨兵 kv=0/score=-inf(KVAndScore.clear 语义),其余行由内核在 # 组起点覆写,无需按页清零。last_dim = 2*coff*head_dim(overlap coff=2)。 - state_rows = self.swa_num_pages * DSV4_C4_STATE_RING + DSV4_C4_STATE_RING + 1 - state_rows = _ceil_div(state_rows, 4) * 4 + state_rows = self._paged_state_rows(self.swa_num_pages, DSV4_C4_STATE_RING, 4) self.c4_state_buffer = torch.zeros( (self.n_c4, state_rows, 4 * self.head_dim), dtype=torch.float32, device="cuda" ) @@ -359,8 +377,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): (self.n_c4, state_rows, 4 * self.indexer_head_dim), dtype=torch.float32, device="cuda" ) for buf in (self.c4_state_buffer, self.c4_indexer_state_buffer): - half = buf.shape[-1] // 2 - buf[:, -1, half:].fill_(float("-inf")) + self._init_state_sentinel(buf) if self.n_c128 > 0: self.c128_pool = PackedPagePool( size=self.c128_size, @@ -375,6 +392,13 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): ) self.full_to_c128_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") self.full_to_c128_indexs[size] = self.c128_pool.HOLD_TOKEN_MEMINDEX + # c128 compressor 在途状态: 与 c4 同样由 full->swa 推导行号,但 ring=128 且无 overlap。 + # last_dim = 2*head_dim;末行是 swa 缺失/出窗时读取的哨兵。 + state_rows = self._paged_state_rows(self.swa_num_pages, DSV4_C128_STATE_RING, 128) + self.c128_state_buffer = torch.zeros( + (self.n_c128, state_rows, 2 * self.head_dim), dtype=torch.float32, device="cuda" + ) + self._init_state_sentinel(self.c128_state_buffer) logger.info( f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " @@ -410,6 +434,10 @@ def get_c4_indexer_state_buffer(self, layer_index: int) -> torch.Tensor: assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 paged indexer state" return self.c4_indexer_state_buffer[self.layer_to_c4_idx[layer_index]] + def get_c128_state_buffer(self, layer_index: int) -> torch.Tensor: + assert self.compress_rates[layer_index] == 128, "只有 c128(HCA) 层有 paged compressor state" + return self.c128_state_buffer[self.layer_to_c128_idx[layer_index]] + # ------------------------------------------------------------------ swa slot lifecycle def set_swa_pressure_valve(self, valve) -> None: """valve(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点回收 swa)。""" diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 24b39b71ce..a0faccb00c 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -29,8 +29,8 @@ class DeepseekV4PromptCachePayload: """prompt cache 载荷: 只剩 swa 按页有效性 bitmap。 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 - 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4 compressor 状态以 swa 页派生 - 寻址(随 swa 页生灭,命中零拷贝续算);c128 状态在 128 边界自然归零,无需恢复。 + 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4/c128 compressor 状态以 swa + 页派生寻址(随 swa 页生灭,命中零拷贝续算)。c128 partial state 不跨 radix 的 128 边界保存。 * ``swa_page_valid``: cpu bool [cache_len // page],插入时按当下 full_to_swa 映射写定 (页内 128 个映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 128 边界, @@ -368,9 +368,8 @@ class DeepseekV4ReqManager(ReqManager): 本类只负责 prep 阶段的分配与 scatter(``prepare_prefill_compress_slots`` / ``prepare_decode_compress_slots``)——必须先于 attention metadata 构建/图捕获; 条目内容由 layer-infer 的 compressor 前向写入。 - * ``req_to_c128_state_pool`` —— c128 compressor 的在途状态(per req、per c128 层)。 - c128 在线聚合在 128 边界自然归零(命中边界必 128 对齐),无缓存常驻需求,保持 req 键控。 - c4 状态(跨边界 overlap)在 mem manager 的 swa 页派生池,随页生灭,命中零拷贝续算。 + * compressor 在途状态不在本类: c4/c128 都在 mem manager 的 swa 页派生池, + 随页生灭,命中零拷贝续算。 * SWA 槽位分配/出窗回收(``prepare_prefill_swa`` / ``prepare_decode_swa``): 每步 prep 阶段 为新 token 调 mem_manager.alloc_swa,并按 per-req 水位线(``_swa_evict_marks``)惰性回收 已出窗位置的 swa 槽。水位线首次置为该请求首个 chunk 的 ready_cache_len(radix 共享前缀 @@ -419,16 +418,6 @@ def __init__( self.layer_to_c128_idx[lid] = c128 c128 += 1 - # c128 compressor 在途状态(fp32): 在线聚合在 128 边界自然归零(命中边界必 128 对齐), - # 无缓存常驻需求,保持 req 键控。c4 状态(有跨边界 overlap)在 mem manager 的 - # swa 页派生池(c4_state_buffer / c4_indexer_state_buffer)。 - self.req_to_c128_state_pool = LayerCache( - size=max_request_num + 1, - dtype=torch.float32, - shape=(1, 128, 2 * head_dim), - layer_num=self.n_c128, - device="cuda", - ) return def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): @@ -524,20 +513,11 @@ def prepare_decode_swa( self.mem_manager.alloc_swa_decode(b_req_idx, b_seq_len, mem_indexes, self.req_to_token_indexs) return - def _reset_state_pool_req(self, cache: LayerCache, req_idx: int): - if cache.layer_num == 0: - return - cache.buffer[:, req_idx, ...].fill_(0) - return - def init_compress_state(self, req_idx: int): - """新请求开始时重置其 compressor 在途状态(对应 mamba 的 init_linear_att_state)。 + """新请求开始时重置 runtime 水位线(对应 mamba 的 init_linear_att_state 调用点)。 - 只有 c128 状态是 req 键控的(c4 状态随 swa 页生灭,内核组起点覆写,无需重置; - 压缩槽位以 full 槽位为键,随请求 full 槽的释放级联回收)。""" + c4/c128 compressor state 都随 swa 页寻址,由内核按组覆写;请求复用时不做 per-req 清零。""" self.clear_runtime_state(req_idx) - if self.n_c128 > 0: - self._reset_state_pool_req(self.req_to_c128_state_pool, req_idx) return # ------------------------------------------------------------------ compress slot prep (per step) @@ -628,14 +608,6 @@ def clear_runtime_state(self, req_idx: int): self._swa_evict_marks[req_idx] = -1 return - def get_compress_state_pool_for_req(self, layer_index: int, req_idx: int): - assert self.compress_rates[layer_index] == 128, "c4 state 在 mem manager 的 swa 页派生池" - return self.req_to_c128_state_pool.buffer[self.layer_to_c128_idx[layer_index], req_idx] - - def get_compress_state_pool(self, layer_index: int): - assert self.compress_rates[layer_index] == 128, "c4 state 在 mem manager 的 swa 页派生池" - return self.req_to_c128_state_pool.buffer[self.layer_to_c128_idx[layer_index]] - def get_prompt_cache_value_ops(self): return DeepseekV4PromptCacheValueOps(self) @@ -712,6 +684,4 @@ def free_req(self, free_req_index: int): def free_all(self): super().free_all() self._swa_evict_marks = [-1 for _ in range(self.max_request_num + 1)] - if self.n_c128 > 0: - self.req_to_c128_state_pool.buffer.fill_(0) return diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 129a93ae0d..4857740992 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -6,6 +6,12 @@ import triton.language as tl from triton.language.extra import libdevice +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_C4_STATE_RING, + DSV4_C128_STATE_RING, + DSV4_SWA_PAGE_SIZE, +) + _SGLANG_COMPRESS_ERR = None _SGLANG_COMPRESS_MOD = None @@ -80,7 +86,7 @@ def _save_partial_states_kernel( IS_C4: tl.constexpr, IS_PREFILL: tl.constexpr, SWA_PAGE_SIZE: tl.constexpr, - C4_STATE_RING: tl.constexpr, + STATE_RING: tl.constexpr, BLOCK: tl.constexpr, ): token_idx = tl.program_id(0) @@ -89,19 +95,18 @@ def _save_partial_states_kernel( seq_len = tl.load(b_seq_len + batch_idx) if IS_C4: - same_page_next = (position % SWA_PAGE_SIZE) + C4_STATE_RING < SWA_PAGE_SIZE - if same_page_next and position + C4_STATE_RING < seq_len: - return - full_slot = tl.load(mem_index + token_idx).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) - if swa_slot < 0: + same_page_next = (position % SWA_PAGE_SIZE) + STATE_RING < SWA_PAGE_SIZE + if same_page_next and position + STATE_RING < seq_len: return - state_row = (swa_slot // SWA_PAGE_SIZE) * C4_STATE_RING + (swa_slot % C4_STATE_RING) else: if position + COMPRESS_RATIO < seq_len: return - req_idx = tl.load(b_req_idx + batch_idx).to(tl.int64) - state_row = req_idx * COMPRESS_RATIO + (position % COMPRESS_RATIO) + + full_slot = tl.load(mem_index + token_idx).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + if swa_slot < 0: + return + state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) offs = tl.arange(0, BLOCK) mask = offs < STATE_WIDTH @@ -144,7 +149,7 @@ def _fused_compress_norm_rope_insert_kernel( IS_C4: tl.constexpr, IS_PREFILL: tl.constexpr, SWA_PAGE_SIZE: tl.constexpr, - C4_STATE_RING: tl.constexpr, + STATE_RING: tl.constexpr, ROPE_HEAD_DIM: tl.constexpr, FP8_MAX: tl.constexpr, SCALE_MIN: tl.constexpr, @@ -188,12 +193,18 @@ def _fused_compress_norm_rope_insert_kernel( other=0, ).to(tl.int64) swa_slot = tl.load(full_to_swa + full_slot, mask=valid_pos & (~use_current), other=-1).to(tl.int64) - state_row = (swa_slot // SWA_PAGE_SIZE) * C4_STATE_RING + (swa_slot % C4_STATE_RING) + state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) state_valid = valid_pos & (~use_current) & (swa_slot >= 0) head_offset = tl.where(token_offsets >= COMPRESS_RATIO, HEAD_DIM, 0) else: - state_row = req_idx * COMPRESS_RATIO + (gather_pos % COMPRESS_RATIO) - state_valid = valid_pos & (~use_current) + full_slot = tl.load( + req_to_token + req_idx * req_to_token_stride0 + gather_pos, + mask=valid_pos & (~use_current), + other=0, + ).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot, mask=valid_pos & (~use_current), other=-1).to(tl.int64) + state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) + state_valid = valid_pos & (~use_current) & (swa_slot >= 0) head_offset = token_offsets * 0 offs = tl.arange(0, BLOCK) @@ -288,7 +299,7 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int) out_pool = mem_manager.c4_pool elif compress_ratio == 128: out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] - state_buffer = infer_state.req_manager.get_compress_state_pool(layer_idx) + state_buffer = mem_manager.get_c128_state_buffer(layer_idx) out_pool = mem_manager.c128_pool else: raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") @@ -367,6 +378,7 @@ def fused_compress( state_width = kv_score.shape[-1] // 2 state_last_dim = metadata.state_buffer.shape[-1] is_c4 = compress_ratio == 4 + state_ring = DSV4_C4_STATE_RING if is_c4 else DSV4_C128_STATE_RING block_state = triton.next_power_of_2(state_width) block_head = triton.next_power_of_2(head_dim) @@ -399,8 +411,8 @@ def fused_compress( WINDOW_SIZE=compress_ratio * (2 if is_c4 else 1), IS_C4=is_c4, IS_PREFILL=metadata.is_prefill, - SWA_PAGE_SIZE=128, - C4_STATE_RING=8, + SWA_PAGE_SIZE=DSV4_SWA_PAGE_SIZE, + STATE_RING=state_ring, ROPE_HEAD_DIM=qk_rope_head_dim, FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, SCALE_MIN=1e-4, @@ -429,8 +441,8 @@ def fused_compress( COMPRESS_RATIO=compress_ratio, IS_C4=is_c4, IS_PREFILL=metadata.is_prefill, - SWA_PAGE_SIZE=128, - C4_STATE_RING=8, + SWA_PAGE_SIZE=DSV4_SWA_PAGE_SIZE, + STATE_RING=state_ring, BLOCK=block_state, num_warps=4, ) @@ -623,7 +635,7 @@ def compressor_decode_step_batch( return out.to(x_new.dtype), should_compress -# ---------------------------------------------------------------------------- paged state (c4) +# ---------------------------------------------------------------------------- paged state # 与 sglang srt compressor 的 paged 路径同构(compress_old 内核 + 分组槽 indices + overlap # extra_data): state 槽位由 swa 槽位算术派生(翻译③ state_loc = page*ring + swa_loc%ring, # 分组槽 = state_loc//ratio),state 随 swa 页生灭,radix 命中零拷贝续算。 @@ -667,35 +679,49 @@ def paged_decode_state_slots( ratio: int, hold_req_id: int, num_swa_pages: int, + overlap: bool = True, ): - """decode 步的 state 分组槽(写槽 = 当前组 clip_down(seq-1) 的槽,overlap 伙伴 = 前一组)。 + """decode 步的 state 分组槽(写槽 = 当前组 clip_down(seq-1) 的槽,可选 overlap 前一组)。 纯张量算术(prep 已写本步 req_to_token),图安全。padding(HOLD)行重定向到 HOLD 页的 state 槽,隔离其垃圾累加。""" seq = b_seq_len.long() write_positions = torch.div(seq - 1, ratio, rounding_mode="floor") * ratio write_slot = _paged_state_group_slot(req_to_token, full_to_swa, b_req_idx, write_positions, page_size, ring, ratio) - overlap_slot = _paged_state_group_slot( - req_to_token, full_to_swa, b_req_idx, write_positions - ratio, page_size, ring, ratio - ) + overlap_slot = None + if overlap: + overlap_slot = _paged_state_group_slot( + req_to_token, full_to_swa, b_req_idx, write_positions - ratio, page_size, ring, ratio + ) hold_slot = num_swa_pages * ring // ratio # HOLD 页区域([pages*ring, pages*ring+ring))的首个分组槽 is_hold = b_req_idx.long() == hold_req_id write_slot = torch.where(is_hold, torch.full_like(write_slot, hold_slot), write_slot) - overlap_slot = torch.where(is_hold, torch.full_like(overlap_slot, hold_slot), overlap_slot) + if overlap_slot is not None: + overlap_slot = torch.where(is_hold, torch.full_like(overlap_slot, hold_slot), overlap_slot) return write_slot, overlap_slot -def paged_prefill_compress_data(req_to_token, full_to_swa, req_idx: int, ready_len: int, seq_len: int, ring: int): +def paged_prefill_compress_data( + req_to_token, + full_to_swa, + req_idx: int, + ready_len: int, + seq_len: int, + ring: int, + ratio: int = 4, + page_size: int = DSV4_SWA_PAGE_SIZE, + overlap: bool = True, +): """单请求 prefill chunk 的 (indices, extra_data, plan): 与 sglang 同走 - triton_create_paged_compress_data(按请求产出,内核经 plan 逐 token 步进)。仅 c4(overlap)。 + triton_create_paged_compress_data(按请求产出,内核经 plan 逐 token 步进)。 三者都与层无关,同一 forward 内可跨全部 c4 层复用。""" mod, _ = _load_sglang_compressor() fn = _load_paged_compress_data_fn() device = req_to_token.device n_new = seq_len - ready_len write_loc, extra_data = fn( - compress_ratio=4, - is_overlap=True, - swa_page_size=128, + compress_ratio=ratio, + is_overlap=overlap, + swa_page_size=page_size, ring_size=ring, req_pool_indices=torch.tensor([req_idx], device=device, dtype=torch.int64), seq_lens=torch.tensor([seq_len], device=device, dtype=torch.int64), @@ -704,7 +730,7 @@ def paged_prefill_compress_data(req_to_token, full_to_swa, req_idx: int, ready_l full_to_swa_index_mapping=full_to_swa, ) plan = mod.CompressorPrefillPlan.generate( - 4, + ratio, n_new, torch.tensor([seq_len], dtype=torch.int64), torch.tensor([n_new], dtype=torch.int64), @@ -727,15 +753,16 @@ def compressor_paged_prefill( compress_data, ready_len, seq_len, + ratio: int = 4, ): - """单请求 prefill/extend chunk(c4 paged): x [n_new, dim] 为位置 [ready, seq) 的 hidden, + """单请求 prefill/extend chunk(paged): x [n_new, dim] 为位置 [ready, seq) 的 hidden, state 写到 swa 派生的分组槽(compress_data 来自 paged_prefill_compress_data,跨层复用)。 - 返回本 chunk 完结组的压缩条目 [seq//4 - ready//4, head_dim](rope 已施加)。""" + 返回本 chunk 完结组的压缩条目 [seq//ratio - ready//ratio, head_dim](rope 已施加)。""" mod, _ = _load_sglang_compressor() - ratio = 4 kv_score = _project_kv_score(x, wkv_w, wgate_w) pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) write_loc, extra_data, plan = compress_data + kwargs = {"extra_data": extra_data} if extra_data is not None else {} out = mod.compress_forward( pool, kv_score, @@ -744,7 +771,7 @@ def compressor_paged_prefill( plan, head_dim=head_dim, compress_ratio=ratio, - extra_data=extra_data, + **kwargs, ) ncomp = seq_len // ratio - ready_len // ratio if ncomp == 0: @@ -768,15 +795,16 @@ def compressor_paged_decode_batch( write_slot, overlap_slot, b_seq_len, + ratio: int = 4, ): - """批量 decode 一步(c4 paged): state 槽位为 swa 派生分组槽(paged_decode_state_slots, + """批量 decode 一步(paged): state 槽位为 swa 派生分组槽(paged_decode_state_slots, 可跨层复用)。返回 (entries [bs, head_dim], should_compress [bs])。""" mod, _ = _load_sglang_compressor() - ratio = 4 kv_score = _project_kv_score(x_new, wkv_w, wgate_w) pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) seq_lens = b_seq_len.to(torch.int32).contiguous() plan = mod.CompressorDecodePlan(ratio, seq_lens) + kwargs = {"extra_data": overlap_slot.view(-1, 1)} if overlap_slot is not None else {} out = mod.compress_forward( pool, kv_score, @@ -785,7 +813,7 @@ def compressor_paged_decode_batch( plan, head_dim=head_dim, compress_ratio=ratio, - extra_data=overlap_slot.view(-1, 1), + **kwargs, ) should_compress = (seq_lens % ratio) == 0 mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) From 07d230865f598d7d5079c757e7d65457e1d389b2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 15 Jun 2026 06:37:11 +0000 Subject: [PATCH 019/214] refact --- .../attention/nsa/fp8_flashmla_sparse.py | 279 ++---------- lightllm/models/deepseek_v4/infer_struct.py | 23 + .../deepseek_v4/layer_infer/compressor.py | 429 ++---------------- .../layer_infer/transformer_layer_infer.py | 247 ++++++++-- .../triton_kernel/build_swa_index_dsv4.py | 75 +++ 5 files changed, 371 insertions(+), 682 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 14b1b3307d..03595c1bce 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -1,5 +1,4 @@ import dataclasses -import inspect import torch from typing import TYPE_CHECKING, Tuple @@ -10,7 +9,6 @@ from lightllm.common.basemodel.infer_struct import InferStateInfo -FLASHMLA_INDEX_ALIGN = 64 # this flash_mla extra-cache fork only instantiates h_q in {64, 128}; pad TP-split q heads up # to the nearest supported count (zero heads are discarded from the output slice). FLASHMLA_SUPPORTED_HEADS = (64, 128) @@ -39,15 +37,6 @@ def _missing_attention_op(feature: str) -> None: ) -def _pad_last_dim(x: torch.Tensor, multiple: int = FLASHMLA_INDEX_ALIGN, value: int = -1) -> torch.Tensor: - pad = (-x.shape[-1]) % multiple - if pad == 0: - return x.contiguous() - out = torch.full((*x.shape[:-1], x.shape[-1] + pad), value, dtype=x.dtype, device=x.device) - out[..., : x.shape[-1]] = x - return out.contiguous() - - def _view_dsv4_flashmla_cache(layer_buffer: torch.Tensor, page_size: int) -> torch.Tensor: from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_MLA_BYTES_PER_TOKEN @@ -55,191 +44,6 @@ def _view_dsv4_flashmla_cache(layer_buffer: torch.Tensor, page_size: int) -> tor return layer_buffer[:, :usable].view(layer_buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) -def _load_flash_mla_with_extra(): - try: - import flash_mla - except Exception as exc: - raise DeepseekV4MissingOperatorError( - "DeepSeek-V4 packed FlashMLA requires the flash_mla package with compiled CUDA extension. " - f"Import failed with: {type(exc).__name__}: {exc}" - ) from exc - - fn = getattr(flash_mla, "flash_mla_with_kvcache", None) - get_mla_metadata = getattr(flash_mla, "get_mla_metadata", None) - missing_symbols = [] - if fn is None: - missing_symbols.append("flash_mla_with_kvcache") - if get_mla_metadata is None: - missing_symbols.append("get_mla_metadata") - if missing_symbols: - raise DeepseekV4MissingOperatorError( - "DeepSeek-V4 requires flash_mla.flash_mla_with_kvcache extra-cache wrapper. " - f"Current module={getattr(flash_mla, '__file__', '')} " - f"is missing symbols {missing_symbols}." - ) - - sig = inspect.signature(fn) - required = { - "attn_sink", - "extra_k_cache", - "extra_indices_in_kvcache", - "topk_length", - "extra_topk_length", - } - missing = sorted(required.difference(sig.parameters)) - if missing: - raise DeepseekV4MissingOperatorError( - "DeepSeek-V4 requires flash_mla.flash_mla_with_kvcache with extra-cache arguments. " - f"Current module={getattr(flash_mla, '__file__', '')} is missing {missing}." - ) - return flash_mla - - -def _build_dsv4_repeated_prefill_reqs(infer_state) -> torch.Tensor: - return torch.repeat_interleave(infer_state.b_req_idx, infer_state.b_q_seq_len.long()) - - -def _build_dsv4_prefill_positions(infer_state) -> torch.Tensor: - total = infer_state.total_token_num - infer_state.prefix_total_token_num - token_offsets = torch.arange(total, dtype=torch.int32, device=infer_state.b_q_seq_len.device) - req_ids = torch.repeat_interleave( - torch.arange(infer_state.batch_size, dtype=torch.long, device=infer_state.b_q_seq_len.device), - infer_state.b_q_seq_len.long(), - ) - local_offsets = token_offsets - infer_state.b_q_start_loc[req_ids] - return infer_state.b_ready_cache_len[req_ids] + local_offsets - - -def _build_dsv4_swa_indices( - req_manager, - mem_manager, - req_idx: torch.Tensor, - positions: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - window = int(mem_manager.sliding_window) - offsets = positions[:, None] - torch.arange(window, dtype=positions.dtype, device=positions.device)[None, :] - valid_pos = offsets >= 0 - safe_offsets = offsets.clamp_min(0).long() - full_slots = req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_offsets] - swa_slots = mem_manager.full_to_swa_indexs[full_slots.long()].to(torch.int32) - indices = torch.where(valid_pos, swa_slots, torch.full_like(swa_slots, -1)) - lengths = torch.clamp(positions + 1, min=1, max=window).to(torch.int32) - return _pad_last_dim(indices.to(torch.int32)).unsqueeze(1), lengths.contiguous() - - -def _gather_dsv4_compress_slots( - infer_state, - mapping: torch.Tensor, - req_idx: torch.Tensor, - valid: torch.Tensor, - offsets: torch.Tensor, - ratio: int, -) -> torch.Tensor: - """条目 g 的压缩槽 = full_to_c*[req_to_token[req, (g+1)*ratio-1]](组末 token 的 full 槽位)。 - 无效条目(超出因果长度/HOLD 行)用位置 0 安全 gather 后由调用方按 valid 掩掉。""" - end_pos = offsets[None, :] * ratio + (ratio - 1) - safe_pos = torch.where(valid, end_pos, torch.zeros_like(end_pos)) - full_slots = infer_state.req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_pos] - return mapping[full_slots.long()].to(torch.int32) - - -def _build_dsv4_c128_indices( - infer_state, - req_idx: torch.Tensor, - positions: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - raw_lengths = (positions + 1) // 128 - lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) - max_len = max(1, int(infer_state.max_kv_seq_len) // 128) - offsets = torch.arange(max_len, dtype=torch.long, device=positions.device) - valid = offsets[None, :] < raw_lengths[:, None] - slots = _gather_dsv4_compress_slots( - infer_state, infer_state.mem_manager.full_to_c128_indexs, req_idx, valid, offsets, 128 - ) - indices = torch.where(valid, slots, torch.full_like(slots, -1)) - return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() - - -def _build_dsv4_c4_indices( - infer_state, - layer_index: int, - req_idx: torch.Tensor, - positions: torch.Tensor, - nsa_dict: dict, -) -> Tuple[torch.Tensor, torch.Tensor]: - """c4(CSA) extra indices: causal all-entries when the entry space fits index_topk, - otherwise Lightning-Indexer scored top-k. Pure tensor ops (decode runs inside cuda graphs).""" - import torch.distributed as dist - import torch.nn.functional as F - from lightllm.distributed.communication_op import all_reduce - - mem_manager = infer_state.mem_manager - raw_lengths = (positions + 1) // 4 - max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) - index_topk = int(nsa_dict["index_topk"]) - offsets = torch.arange(max_entries, dtype=torch.long, device=positions.device) - valid = offsets[None, :] < raw_lengths[:, None] - slots = _gather_dsv4_compress_slots(infer_state, mem_manager.full_to_c4_indexs, req_idx, valid, offsets, 4) - - if max_entries <= index_topk: - lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) - indices = torch.where(valid, slots, torch.full_like(slots, -1)) - return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() - - idx_q = nsa_dict["idx_q"] # [T, H, index_head_dim], rope applied - idx_weight = nsa_dict["idx_weight"] # [T, H] fp32, weight scale applied - score_scale = float(nsa_dict["indexer_score_scale"]) - hold_slot = mem_manager.c4_indexer_pool.HOLD_TOKEN_MEMINDEX - safe_slots = torch.where(valid, slots.long(), torch.full_like(slots.long(), hold_slot)) - k = mem_manager.gather_indexer_k(layer_index, safe_slots.reshape(-1)).view(positions.shape[0], max_entries, -1) - - num_tokens, num_heads = idx_q.shape[0], idx_q.shape[1] - score_chunks = [] - chunk = max(1, min(num_tokens, (16 * 1024 * 1024) // max(1, num_heads * max_entries))) - for start in range(0, num_tokens, chunk): - end = min(num_tokens, start + chunk) - scores = torch.einsum("thd,tnd->thn", idx_q[start:end].float(), k[start:end].float()) - scores = F.relu(scores) * score_scale - score_chunks.append((scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1)) - index_scores = torch.cat(score_chunks, dim=0) - if int(nsa_dict.get("tp_world_size", 1)) > 1: - all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - index_scores = index_scores.masked_fill(~valid, float("-inf")) - top = index_scores.topk(index_topk, dim=-1).indices - top_valid = torch.gather(valid, 1, top) - top_slots = torch.gather(slots.long(), 1, top).to(torch.int32) - indices = torch.where(top_valid, top_slots, torch.full_like(top_slots, -1)) - lengths = torch.clamp(torch.minimum(raw_lengths, torch.full_like(raw_lengths, index_topk)), min=1) - return _pad_last_dim(indices).unsqueeze(1), lengths.to(torch.int32).contiguous() - - -def _build_dsv4_extra_metadata( - infer_state, - layer_index: int, - compress_ratio: int, - req_idx: torch.Tensor, - positions: torch.Tensor, - swa_indices: torch.Tensor, - swa_lengths: torch.Tensor, - nsa_dict: dict, -) -> "_Dsv4Metadata": - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C128_PAGE_SIZE, DSV4_C4_PAGE_SIZE - - if compress_ratio == 0: - return _Dsv4Metadata(swa_indices, swa_lengths) - if compress_ratio == 4: - extra_indices, extra_lengths = _build_dsv4_c4_indices(infer_state, layer_index, req_idx, positions, nsa_dict) - extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(layer_index) - extra_cache = _view_dsv4_flashmla_cache(extra_buffer, DSV4_C4_PAGE_SIZE) - return _Dsv4Metadata(swa_indices, swa_lengths, extra_cache, extra_indices, extra_lengths) - if compress_ratio == 128: - extra_indices, extra_lengths = _build_dsv4_c128_indices(infer_state, req_idx, positions) - extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(layer_index) - extra_cache = _view_dsv4_flashmla_cache(extra_buffer, DSV4_C128_PAGE_SIZE) - return _Dsv4Metadata(swa_indices, swa_lengths, extra_cache, extra_indices, extra_lengths) - raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") - - @dataclasses.dataclass class _Dsv4Metadata: swa_indices: torch.Tensor @@ -249,6 +53,28 @@ class _Dsv4Metadata: extra_lengths: torch.Tensor = None +def _metadata_from_dict(infer_state, nsa_dict: dict) -> "_Dsv4Metadata": + """Bundle the model-built FINAL index tensors (carried in nsa_dict by DeepseekV4IndexInfer) with + the layer-keyed fp8 extra-cache byte view. The cache view is data-independent (a fixed per-layer + buffer slice), so it is built here -- a genuine flash_mla ABI concern -- rather than on the model + side; only the index/length tensors cross the att_control boundary.""" + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C128_PAGE_SIZE, DSV4_C4_PAGE_SIZE + + ratio = nsa_dict["compress_ratio"] + extra_cache = None + if ratio: + page = DSV4_C4_PAGE_SIZE if ratio == 4 else DSV4_C128_PAGE_SIZE + extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(nsa_dict["layer_index"]) + extra_cache = _view_dsv4_flashmla_cache(extra_buffer, page) + return _Dsv4Metadata( + swa_indices=nsa_dict["swa_indices"], + swa_lengths=nsa_dict["swa_lengths"], + extra_cache=extra_cache, + extra_indices=nsa_dict.get("extra_indices"), + extra_lengths=nsa_dict.get("extra_lengths"), + ) + + class NsaFlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) @@ -257,7 +83,6 @@ def __init__(self, model): torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) for _ in range(2) ] - self._flash_mla = None def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparsePrefillAttState": return NsaFlashMlaFp8SparsePrefillAttState(backend=self, infer_state=infer_state) @@ -265,11 +90,6 @@ def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMl def create_att_decode_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparseDecodeAttState": return NsaFlashMlaFp8SparseDecodeAttState(backend=self, infer_state=infer_state) - def flash_mla(self): - if self._flash_mla is None: - self._flash_mla = _load_flash_mla_with_extra() - return self._flash_mla - @dataclasses.dataclass class NsaFlashMlaFp8SparsePrefillAttState(BasePrefillAttState): @@ -360,30 +180,9 @@ def _nsa_prefill_att( ) return mla_out - def _build_flashmla_kvcache_prefill_metadata(self, nsa_dict: dict) -> _Dsv4Metadata: - infer_state = self.infer_state - req_idx = _build_dsv4_repeated_prefill_reqs(infer_state) - positions = _build_dsv4_prefill_positions(infer_state) - swa_indices, swa_lengths = _build_dsv4_swa_indices( - infer_state.req_manager, - infer_state.mem_manager, - req_idx, - positions, - ) - return _build_dsv4_extra_metadata( - infer_state, - nsa_dict["layer_index"], - nsa_dict["compress_ratio"], - req_idx, - positions, - swa_indices, - swa_lengths, - nsa_dict, - ) - def _flashmla_kvcache_prefill_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() - metadata = self._build_flashmla_kvcache_prefill_metadata(nsa_dict) + metadata = _metadata_from_dict(self.infer_state, nsa_dict) return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) def _flashmla_kvcache_att( @@ -394,7 +193,7 @@ def _flashmla_kvcache_att( attn_sink: torch.Tensor, nsa_dict: dict, ) -> torch.Tensor: - flash_mla = self.backend.flash_mla() + import flash_mla from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE q_4d = q.unsqueeze(1).contiguous() @@ -458,7 +257,8 @@ def init_state(self): ragged_mem_index=self.ragged_mem_index, hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, ) - flash_mla = self.backend.flash_mla() + import flash_mla + # one sched_meta per layer type: the lazy config locks extra-cache geometry (page size, # presence) on first invocation, so swa-only/c4/c128 layers must not share one object. self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} @@ -468,7 +268,8 @@ def reset_sched_meta_for_capture(self): # cuda-graph capture hook: the warmup pass already locked/stored sched meta on this # (shared) state object; reset so the capture pass re-plans INSIDE the graph and every # replay re-plans from the live tensors instead of binding warmup leftovers. - flash_mla = self.backend.flash_mla() + import flash_mla + self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} return @@ -539,29 +340,9 @@ def _nsa_decode_att( ) return o_tensor[:, 0, :, :] # [b, 1, h, d] -> [b, h, d] - def _build_flashmla_kvcache_decode_metadata(self, nsa_dict: dict) -> _Dsv4Metadata: - infer_state = self.infer_state - positions = infer_state.b_seq_len.to(torch.int32) - 1 - swa_indices, swa_lengths = _build_dsv4_swa_indices( - infer_state.req_manager, - infer_state.mem_manager, - infer_state.b_req_idx, - positions, - ) - return _build_dsv4_extra_metadata( - infer_state, - nsa_dict["layer_index"], - nsa_dict["compress_ratio"], - infer_state.b_req_idx, - positions, - swa_indices, - swa_lengths, - nsa_dict, - ) - def _flashmla_kvcache_decode_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() - metadata = self._build_flashmla_kvcache_decode_metadata(nsa_dict) + metadata = _metadata_from_dict(self.infer_state, nsa_dict) return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) def _flashmla_kvcache_att( @@ -572,7 +353,7 @@ def _flashmla_kvcache_att( attn_sink: torch.Tensor, nsa_dict: dict, ) -> torch.Tensor: - flash_mla = self.backend.flash_mla() + import flash_mla from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE q_4d = q.unsqueeze(1).contiguous() diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index caf8ca6fa8..ca2ac83b03 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -18,6 +18,11 @@ def __init__(self): self.position_sin_sliding = None self.position_cos_compress = None self.position_sin_compress = None + # layer-independent sparse-index metadata, built once per forward in init_some_extra_state + # (None until then so copy_for_cuda_graph's tensor-attr loop skips them). + self.dsv4_sparse_req_idx = None + self.dsv4_swa_indices = None + self.dsv4_swa_lengths = None def init_some_extra_state(self, model): super().init_some_extra_state(model) # sets position_ids, b_q_seq_len, b_q_start_loc (prefill) @@ -26,6 +31,24 @@ def init_some_extra_state(self, model): self.position_sin_sliding = torch.index_select(model._sin_cached_sliding, 0, pos) self.position_cos_compress = torch.index_select(model._cos_cached_compress, 0, pos) self.position_sin_compress = torch.index_select(model._sin_cached_compress, 0, pos) + # Per-token request id (decode: one token per req; prefill: ragged -> repeat by q-len). + # Layer-independent; the swa kernel + build_metadata's c4/c128 readers all reuse it. + if self.is_prefill: + self.dsv4_sparse_req_idx = torch.repeat_interleave(self.b_req_idx, self.b_q_seq_len.long()) + else: + self.dsv4_sparse_req_idx = self.b_req_idx + # Sliding-window indices: layer-independent (full_to_swa is global, window is const), so build + # once here via one fused kernel instead of recomputing per layer. const [T, window] shape is + # cuda-graph-safe (no max_kv_seq_len dependence) and auto-staged by copy_for_cuda_graph. + from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_swa_index + + self.dsv4_swa_indices, self.dsv4_swa_lengths = build_swa_index( + req_idx=self.dsv4_sparse_req_idx, + positions=self.position_ids, + req_to_token_indexs=self.req_manager.req_to_token_indexs, + full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, + window=int(self.mem_manager.sliding_window), + ) # prefill-cudagraph 桶填充的 HOLD 尾请求的 q 行数。其注意力读 HOLD 槽位(内容被并发写 # 竞争,每轮不同),输出必须清零,否则 pad 行 hidden 不确定 -> MoE 路由抖动 -> 共享 expert # 批次组成变化 -> 真实行 GEMM 归约顺序变化(ulp 级),44 层放大后翻转低置信 token。 diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 4857740992..c66fd22de5 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -6,6 +6,7 @@ import triton.language as tl from triton.language.extra import libdevice +from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( DSV4_C4_STATE_RING, DSV4_C128_STATE_RING, @@ -13,12 +14,6 @@ ) -_SGLANG_COMPRESS_ERR = None -_SGLANG_COMPRESS_MOD = None -_SGLANG_LINEAR_BF16_FP32 = None -_FREQ_CIS_CACHE = {} - - @dataclass class CoreCompressorMetadata: layer_idx: int @@ -159,6 +154,7 @@ def _fused_compress_norm_rope_insert_kernel( PAGE_SIZE: tl.constexpr, BYTES_PER_PAGE: tl.constexpr, BLOCK: tl.constexpr, + OUTPUT_BF16: tl.constexpr, ): token_idx = tl.program_id(0) out_slot = tl.load(out_slots + token_idx).to(tl.int64) @@ -260,6 +256,13 @@ def _fused_compress_norm_rope_insert_kernel( new_odd = odd * cos_v + even * sin_v rotated = tl.interleave(new_even, new_odd) + if OUTPUT_BF16: + # indexer-K path: emit the post-rope full HEAD_DIM vector as dense bf16 (token-indexed), + # leaving the fp8 single-amax pack to destindex_copy_indexer_k_dsv4 (the c4_indexer_pool + # ABI differs from the latent slab: whole-vector fp8 + one fp32 scale, no bf16 rope tail). + tl.store(out_buffer + token_idx * HEAD_DIM + offs, rotated.to(tl.bfloat16), mask=dim_mask) + return + page = out_slot // PAGE_SIZE token_in_page = out_slot % PAGE_SIZE data_base = page * BYTES_PER_PAGE + token_in_page * (NOPE_DIM + ROPE_HEAD_DIM * 2) @@ -288,21 +291,38 @@ def _fused_compress_norm_rope_insert_kernel( return -def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int): +def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, is_in_indexer: bool = False): if compress_ratio == 0: return None - mem_manager = infer_state.mem_manager - if compress_ratio == 4: + mem_manager: DeepseekV4MemoryManager = infer_state.mem_manager + if is_in_indexer: + # c4 Lightning-Indexer key compression: same window/state machinery as the c4 latent + # compressor but with index_head_dim, a separate state pool, and a DENSE bf16 scratch + # out_buffer (the kernel's OUTPUT_BF16 path); the fp8 pack into c4_indexer_pool is done + # afterwards by pack_indexer_k_to_cache. + assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] - state_buffer = mem_manager.get_c4_state_buffer(layer_idx) - out_pool = mem_manager.c4_pool - elif compress_ratio == 128: - out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] - state_buffer = mem_manager.get_c128_state_buffer(layer_idx) - out_pool = mem_manager.c128_pool + state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) + out_buffer = torch.empty( + (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), + dtype=torch.bfloat16, + device=infer_state.mem_index.device, + ) + out_page_size = 1 # unused under OUTPUT_BF16 (token-indexed dense scratch, not paged) else: - raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") + if compress_ratio == 4: + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + state_buffer = mem_manager.get_c4_state_buffer(layer_idx) + out_pool = mem_manager.c4_pool + elif compress_ratio == 128: + out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] + state_buffer = mem_manager.get_c128_state_buffer(layer_idx) + out_pool = mem_manager.c128_pool + else: + raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") + out_buffer = mem_manager.get_compressed_kv_buffer(layer_idx) + out_page_size = out_pool.page_size token_to_batch_idx = infer_state.b_req_idx if infer_state.is_prefill: @@ -319,8 +339,8 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int) out_slots=out_slots, mem_index=infer_state.mem_index, state_buffer=state_buffer, - out_buffer=mem_manager.get_compressed_kv_buffer(layer_idx), - out_page_size=out_pool.page_size, + out_buffer=out_buffer, + out_page_size=out_page_size, position_ids=infer_state.position_ids, b_req_idx=infer_state.b_req_idx, b_seq_len=infer_state.b_seq_len, @@ -371,6 +391,7 @@ def fused_compress( compress_ratio: int, cos_table: torch.Tensor, sin_table: torch.Tensor, + output_bf16: bool = False, ): if metadata is None or kv_score.shape[0] == 0: return @@ -422,6 +443,7 @@ def fused_compress( PAGE_SIZE=metadata.out_page_size, BYTES_PER_PAGE=metadata.out_buffer.shape[-1], BLOCK=block_head, + OUTPUT_BF16=output_bf16, num_warps=4, ) @@ -447,374 +469,3 @@ def fused_compress( num_warps=4, ) return - - -def _load_sglang_compressor(): - global _SGLANG_COMPRESS_ERR, _SGLANG_COMPRESS_MOD, _SGLANG_LINEAR_BF16_FP32 - if _SGLANG_COMPRESS_MOD is not None: - return _SGLANG_COMPRESS_MOD, _SGLANG_LINEAR_BF16_FP32 - if _SGLANG_COMPRESS_ERR is not None: - raise _SGLANG_COMPRESS_ERR - try: - from sglang.jit_kernel.dsv4 import linear_bf16_fp32 - from sglang.jit_kernel.dsv4 import compress_old as compress_mod - except Exception as exc: - _SGLANG_COMPRESS_ERR = RuntimeError( - "DeepSeek-V4 fused compressor requires sglang.jit_kernel.dsv4 " - "(linear_bf16_fp32 + compress_old). Install/export the SGLang package " - "or vendor the DSv4 compressor JIT into LightLLM." - ) - raise _SGLANG_COMPRESS_ERR from exc - _SGLANG_COMPRESS_MOD = compress_mod - _SGLANG_LINEAR_BF16_FP32 = linear_bf16_fp32 - return compress_mod, linear_bf16_fp32 - - -def _load_paged_compress_data_fn(): - from sglang.jit_kernel.dsv4 import triton_create_paged_compress_data - - return triton_create_paged_compress_data - - -def _freq_cis(cos_table, sin_table): - key = ( - cos_table.data_ptr(), - sin_table.data_ptr(), - cos_table.device, - tuple(cos_table.shape), - tuple(sin_table.shape), - ) - cached = _FREQ_CIS_CACHE.get(key) - if cached is None: - cached = torch.complex(cos_table.float(), sin_table.float()) - _FREQ_CIS_CACHE[key] = cached - return cached - - -def _sglang_ape(ape, ratio, head_dim): - if ratio == 4: - return torch.cat([ape[:, :head_dim], ape[:, head_dim:]], dim=0).contiguous() - return ape.contiguous() - - -def _compressor_weight(wkv_w, wgate_w): - return torch.cat([wkv_w, wgate_w], dim=0).contiguous() - - -def _project_kv_score(x, wkv_w, wgate_w): - _, linear_bf16_fp32 = _load_sglang_compressor() - return linear_bf16_fp32(x, _compressor_weight(wkv_w, wgate_w)) - - -def _state_pool_view(state_pool): - if state_pool is None: - raise RuntimeError("DeepSeek-V4 fused compressor requires a persistent state_pool") - if state_pool.dim() == 4 and state_pool.shape[1] == 1: - return state_pool.squeeze(1) - return state_pool - - -def compressor_prefill_state( - x, - wkv_w, - wgate_w, - norm_w, - ape, - ratio, - head_dim, - cos_table, - sin_table, - eps, - state_pool, -): - """start_pos==0 prefill for ONE request: x [s, dim] -> compressed entries [s//ratio, head_dim] - (rope applied). state_pool is the request's persistent jit state slice [1, slots, coff*2*head_dim]; - it is rebuilt in place so the decode path can continue from the trailing partial window.""" - mod, _ = _load_sglang_compressor() - kv_score = _project_kv_score(x, wkv_w, wgate_w) - pool = _state_pool_view(state_pool) - pool.zero_() - seq_len = x.shape[0] - plan = mod.CompressorPrefillPlan.generate( - ratio, - seq_len, - torch.tensor([seq_len], dtype=torch.int64), - torch.tensor([seq_len], dtype=torch.int64), - x.device, - ) - indices = torch.zeros((1,), device=x.device, dtype=torch.int32) - out = mod.compress_forward( - pool, - kv_score, - _sglang_ape(ape.float(), ratio, head_dim), - indices, - plan, - head_dim=head_dim, - compress_ratio=ratio, - ) - ncomp = seq_len // ratio - if ncomp == 0: - return x.new_zeros(0, head_dim) - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() - return out.index_select(0, ragged_ids).to(x.dtype) - - -def compressor_decode_step_single( - x_new, - wkv_w, - wgate_w, - norm_w, - ape, - ratio, - head_dim, - cos_table, - sin_table, - eps, - state_pool, - start_pos, -): - """One token for ONE request (chunked-prefill extend path). Returns the finished compressed - entry [head_dim] when (start_pos+1) % ratio == 0, else None. Mutates state_pool in place.""" - mod, _ = _load_sglang_compressor() - kv_score = _project_kv_score(x_new.view(1, -1), wkv_w, wgate_w) - pool = _state_pool_view(state_pool) - seq_len = start_pos + 1 - plan = mod.CompressorDecodePlan( - ratio, - torch.tensor([seq_len], device=x_new.device, dtype=torch.int32), - ) - indices = torch.zeros((1,), device=x_new.device, dtype=torch.int32) - out = mod.compress_forward( - pool, - kv_score, - _sglang_ape(ape.float(), ratio, head_dim), - indices, - plan, - head_dim=head_dim, - compress_ratio=ratio, - ) - if seq_len % ratio != 0: - return None - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - return out[0].to(x_new.dtype) - - -def compressor_decode_step_batch( - x_new, - wkv_w, - wgate_w, - norm_w, - ape, - ratio, - head_dim, - rope_dim, - cos_table, - sin_table, - eps, - state_pool, - b_req_idx, - start_pos, -): - mod, _ = _load_sglang_compressor() - kv_score = _project_kv_score(x_new, wkv_w, wgate_w) - pool = _state_pool_view(state_pool) - seq_lens = (start_pos + 1).to(torch.int32).contiguous() - plan = mod.CompressorDecodePlan(ratio, seq_lens) - out = mod.compress_forward( - pool, - kv_score, - _sglang_ape(ape.float(), ratio, head_dim), - b_req_idx.to(torch.int32).contiguous(), - plan, - head_dim=head_dim, - compress_ratio=ratio, - ) - should_compress = (seq_lens % ratio) == 0 - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - return out.to(x_new.dtype), should_compress - - -# ---------------------------------------------------------------------------- paged state -# 与 sglang srt compressor 的 paged 路径同构(compress_old 内核 + 分组槽 indices + overlap -# extra_data): state 槽位由 swa 槽位算术派生(翻译③ state_loc = page*ring + swa_loc%ring, -# 分组槽 = state_loc//ratio),state 随 swa 页生灭,radix 命中零拷贝续算。 - - -def paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: - """state 池行数 = 页数*ring + ring(HOLD 页) + 1(哨兵行),向上取整到 ratio 整除 - (分组视图 [-1, ratio, last_dim] 需要)。与 sglang CompressStatePool 的 _size 公式一致。""" - rows = num_swa_pages * ring + ring + 1 - return (rows + ratio - 1) // ratio * ratio - - -def init_paged_state_pool(buffer: torch.Tensor) -> None: - """末行为哨兵: kv 半边置 0、score 半边置 -inf(KVAndScore.clear 语义)。其余行无需初始化 - (内核在组起点覆写)。buffer: [rows, 2*coff*head_dim] fp32。""" - half = buffer.shape[-1] // 2 - buffer[-1, :half].zero_() - buffer[-1, half:].fill_(float("-inf")) - return - - -def _paged_state_group_slot(req_to_token, full_to_swa, b_req_idx, positions, page_size, ring, ratio): - """位置 -> state 分组槽(= sglang create_paged_compressor_data.get_raw_loc): - state_loc = (swa_loc//page)*ring + swa_loc%ring; 分组槽 = state_loc//ratio。 - 负位置按 sglang 语义 mask 到 0;已出窗(swa_loc<0)的位置落到 -1(哨兵行,score=-inf)。""" - positions = positions.masked_fill(positions < 0, 0) - full = req_to_token[b_req_idx.long(), positions] - swa_loc = full_to_swa[full.long()].long() - state_loc = torch.div(swa_loc, page_size, rounding_mode="floor") * ring + swa_loc % ring - state_loc = torch.where(swa_loc < 0, torch.full_like(state_loc, -1), state_loc) - return torch.div(state_loc, ratio, rounding_mode="floor").to(torch.int32) - - -def paged_decode_state_slots( - req_to_token, - full_to_swa, - b_req_idx, - b_seq_len, - page_size: int, - ring: int, - ratio: int, - hold_req_id: int, - num_swa_pages: int, - overlap: bool = True, -): - """decode 步的 state 分组槽(写槽 = 当前组 clip_down(seq-1) 的槽,可选 overlap 前一组)。 - 纯张量算术(prep 已写本步 req_to_token),图安全。padding(HOLD)行重定向到 HOLD 页的 - state 槽,隔离其垃圾累加。""" - seq = b_seq_len.long() - write_positions = torch.div(seq - 1, ratio, rounding_mode="floor") * ratio - write_slot = _paged_state_group_slot(req_to_token, full_to_swa, b_req_idx, write_positions, page_size, ring, ratio) - overlap_slot = None - if overlap: - overlap_slot = _paged_state_group_slot( - req_to_token, full_to_swa, b_req_idx, write_positions - ratio, page_size, ring, ratio - ) - hold_slot = num_swa_pages * ring // ratio # HOLD 页区域([pages*ring, pages*ring+ring))的首个分组槽 - is_hold = b_req_idx.long() == hold_req_id - write_slot = torch.where(is_hold, torch.full_like(write_slot, hold_slot), write_slot) - if overlap_slot is not None: - overlap_slot = torch.where(is_hold, torch.full_like(overlap_slot, hold_slot), overlap_slot) - return write_slot, overlap_slot - - -def paged_prefill_compress_data( - req_to_token, - full_to_swa, - req_idx: int, - ready_len: int, - seq_len: int, - ring: int, - ratio: int = 4, - page_size: int = DSV4_SWA_PAGE_SIZE, - overlap: bool = True, -): - """单请求 prefill chunk 的 (indices, extra_data, plan): 与 sglang 同走 - triton_create_paged_compress_data(按请求产出,内核经 plan 逐 token 步进)。 - 三者都与层无关,同一 forward 内可跨全部 c4 层复用。""" - mod, _ = _load_sglang_compressor() - fn = _load_paged_compress_data_fn() - device = req_to_token.device - n_new = seq_len - ready_len - write_loc, extra_data = fn( - compress_ratio=ratio, - is_overlap=overlap, - swa_page_size=page_size, - ring_size=ring, - req_pool_indices=torch.tensor([req_idx], device=device, dtype=torch.int64), - seq_lens=torch.tensor([seq_len], device=device, dtype=torch.int64), - extend_seq_lens=torch.tensor([n_new], device=device, dtype=torch.int64), - req_to_token=req_to_token, - full_to_swa_index_mapping=full_to_swa, - ) - plan = mod.CompressorPrefillPlan.generate( - ratio, - n_new, - torch.tensor([seq_len], dtype=torch.int64), - torch.tensor([n_new], dtype=torch.int64), - device, - ) - return write_loc, extra_data, plan - - -def compressor_paged_prefill( - x, - wkv_w, - wgate_w, - norm_w, - ape, - head_dim, - cos_table, - sin_table, - eps, - state_buffer, - compress_data, - ready_len, - seq_len, - ratio: int = 4, -): - """单请求 prefill/extend chunk(paged): x [n_new, dim] 为位置 [ready, seq) 的 hidden, - state 写到 swa 派生的分组槽(compress_data 来自 paged_prefill_compress_data,跨层复用)。 - 返回本 chunk 完结组的压缩条目 [seq//ratio - ready//ratio, head_dim](rope 已施加)。""" - mod, _ = _load_sglang_compressor() - kv_score = _project_kv_score(x, wkv_w, wgate_w) - pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) - write_loc, extra_data, plan = compress_data - kwargs = {"extra_data": extra_data} if extra_data is not None else {} - out = mod.compress_forward( - pool, - kv_score, - _sglang_ape(ape.float(), ratio, head_dim), - write_loc, - plan, - head_dim=head_dim, - compress_ratio=ratio, - **kwargs, - ) - ncomp = seq_len // ratio - ready_len // ratio - if ncomp == 0: - return x.new_zeros(0, head_dim) - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - ragged_ids = plan.compress_plan.view(torch.int32)[:ncomp, 0].long() - return out.index_select(0, ragged_ids).to(x.dtype) - - -def compressor_paged_decode_batch( - x_new, - wkv_w, - wgate_w, - norm_w, - ape, - head_dim, - cos_table, - sin_table, - eps, - state_buffer, - write_slot, - overlap_slot, - b_seq_len, - ratio: int = 4, -): - """批量 decode 一步(paged): state 槽位为 swa 派生分组槽(paged_decode_state_slots, - 可跨层复用)。返回 (entries [bs, head_dim], should_compress [bs])。""" - mod, _ = _load_sglang_compressor() - kv_score = _project_kv_score(x_new, wkv_w, wgate_w) - pool = state_buffer.view(-1, ratio, state_buffer.shape[-1]) - seq_lens = b_seq_len.to(torch.int32).contiguous() - plan = mod.CompressorDecodePlan(ratio, seq_lens) - kwargs = {"extra_data": overlap_slot.view(-1, 1)} if overlap_slot is not None else {} - out = mod.compress_forward( - pool, - kv_score, - _sglang_ape(ape.float(), ratio, head_dim), - write_slot, - plan, - head_dim=head_dim, - compress_ratio=ratio, - **kwargs, - ) - should_compress = (seq_lens % ratio) == 0 - mod.compress_fused_norm_rope_inplace(out, norm_w.float(), eps, _freq_cis(cos_table, sin_table), plan) - return out.to(x_new.dtype), should_compress diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 1b5c5f4c3f..f99df4b00d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -8,6 +8,7 @@ from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor +from lightllm.utils.vllm_utils import vllm_ops from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post from .compressor import fused_compress as fused_compress_op from .compressor import prepare_partial_states @@ -26,9 +27,6 @@ def __init__(self, layer_num, network_config): self.qk_rope_head_dim = network_config["qk_rope_head_dim"] self.qk_nope_head_dim = self.head_dim_ - self.qk_rope_head_dim self.v_head_dim = self.head_dim_ - self.index_n_heads = network_config["index_n_heads"] - self.index_head_dim = network_config["index_head_dim"] - self.index_topk = network_config["index_topk"] self.o_groups = network_config["o_groups"] self.hc_mult = network_config["hc_mult"] self.sinkhorn_iters = network_config["hc_sinkhorn_iters"] @@ -48,14 +46,14 @@ def __init__(self, layer_num, network_config): self.swiglu_limit = network_config["swiglu_limit"] self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) self.tp_q_head_num_ = self.num_heads // self.tp_world_size_ - self.tp_index_n_heads = self.index_n_heads // self.tp_world_size_ self.tp_groups = self.o_groups // self.tp_world_size_ self.enable_ep_moe = get_env_start_args().enable_ep_moe - self.indexer_score_scale = self.index_head_dim ** -0.5 - self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 self.compressor = CompressorInfer( layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ ) + self.index_infer = DeepseekV4IndexInfer( + layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ + ) # ------------------------------------------------------------------ forward (HC-threaded) def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): @@ -180,19 +178,6 @@ def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: Deepsee o = layer_weight.wo_b_.mm(o) return self._tpsp_reduce(input=o, infer_state=infer_state) - # ------------------------------------------------------------------ compressor / indexer - def _indexer_q_weight( - self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - if self.compress_ratio != 4: - return None, None - cos_tok = infer_state.position_cos_compress - sin_tok = infer_state.position_sin_compress - idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_n_heads, self.index_head_dim) - rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) - idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale - return idx_q, idx_weight - # ------------------------------------------------------------------ attention (prefill) def context_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight @@ -241,7 +226,14 @@ def _context_attention_kernel( ): self.compressor.prepare_states(x, infer_state, layer_weight) self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) - idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + # Write this step's c4 Lightning-Indexer keys (no-op off c4) BEFORE build_metadata so the + # scorer's gather_indexer_k reads fresh+accumulated entries. + self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) + # Build the FINAL flash_mla index tensors here (model side), so att_control is a thin + # transport of ready-to-forward tensors -- not indexer raw material. Must stay after + # fused_compress (c4 reads the indexer-K pool it writes) and before prefill_att (keeps the + # c4 einsum/topk/all_reduce at the same cuda-graph capture position). + meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) att_control = AttControl( nsa_prefill=True, nsa_prefill_dict={ @@ -250,14 +242,8 @@ def _context_attention_kernel( "compress_ratio": self.compress_ratio, "head_dim_v": self.v_head_dim, "softmax_scale": self.softmax_scale, - "q_lora": q_lora, - "hidden_states": x, "attn_sink": layer_weight.attn_sink_.weight, - "idx_q": idx_q, - "idx_weight": idx_weight, - "index_topk": self.index_topk, - "indexer_score_scale": self.indexer_score_scale, - "tp_world_size": self.tp_world_size_, + **meta, }, ) out = infer_state.prefill_att_state.prefill_att( @@ -285,7 +271,8 @@ def _token_attention_kernel( ): self.compressor.prepare_states(x, infer_state, layer_weight) self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) - idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) + meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) att_control = AttControl( nsa_decode=True, nsa_decode_dict={ @@ -294,14 +281,8 @@ def _token_attention_kernel( "compress_ratio": self.compress_ratio, "head_dim_v": self.v_head_dim, "softmax_scale": self.softmax_scale, - "q_lora": q_lora, - "hidden_states": x, "attn_sink": layer_weight.attn_sink_.weight, - "idx_q": idx_q, - "idx_weight": idx_weight, - "index_topk": self.index_topk, - "indexer_score_scale": self.indexer_score_scale, - "tp_world_size": self.tp_world_size_, + **meta, }, ) return infer_state.decode_att_state.decode_att( @@ -354,8 +335,6 @@ def _select_experts( def _select_experts_vllm( self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - from vllm import _custom_ops as ops - M = logits.shape[0] bias = None input_tokens = None @@ -373,7 +352,7 @@ def _select_experts_vllm( weights = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.float32, device=logits.device) indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=indices_dtype, device=logits.device) token_expert_indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.int32, device=logits.device) - ops.topk_hash_softplus_sqrt( + vllm_ops.topk_hash_softplus_sqrt( weights, indices, token_expert_indices, @@ -388,11 +367,18 @@ def _select_experts_vllm( class CompressorInfer: - def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): + """Window-softmax compressor. is_in_indexer=False compresses the c4/c128 latent KV into the + paged fp8 slab (attention extra_k); is_in_indexer=True reuses the SAME machinery (mirroring + sglang's Compressor(is_in_indexer=...)) with the indexer weights/dims/state pool to produce the + per-c4-entry Lightning-Indexer keys, emitted as dense bf16 (OUTPUT_BF16) then fp8-packed into + c4_indexer_pool by the caller. Indexer mode is c4-only.""" + + def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int, is_in_indexer: bool = False): super().__init__() self.layer_idx_ = layer_idx self.network_config_ = network_config self.tp_world_size_ = tp_world_size + self.is_in_indexer = is_in_indexer self.compress_ratio = network_config["compress_ratios"][layer_idx] self.head_dim = network_config["head_dim"] self.index_head_dim = network_config["index_head_dim"] @@ -410,13 +396,23 @@ def prepare_states( infer_state=infer_state, layer_idx=self.layer_idx_, compress_ratio=self.compress_ratio, + is_in_indexer=self.is_in_indexer, ) if self._metadata is not None: - self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm(x).float() + if self.is_in_indexer: + # indexer wkv/wgate are two separate replicated weights; cat -> [T, 2*coff*idx_hd] + # (same [kv | score] layout the fused compressor_wkv_gate_ produces for attention). + kv = layer_weight.idx_cmp_wkv_.mm(x) + gate = layer_weight.idx_cmp_wgate_.mm(x) + self._metadata.kv_score = torch.cat([kv, gate], dim=-1).float() + ape = layer_weight.idx_cmp_ape_.weight + else: + self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm(x).float() + ape = layer_weight.compressor_ape_.weight prepare_partial_states( kv_score=self._metadata.kv_score, metadata=self._metadata, - ape=layer_weight.compressor_ape_.weight, + ape=ape, compress_ratio=self.compress_ratio, ) return self._metadata @@ -433,15 +429,178 @@ def fused_compress( metadata = self._metadata if metadata is None: raise RuntimeError("DeepSeek-V4 compressor.prepare_states must run before fused_compress") + if self.is_in_indexer: + norm_weight = layer_weight.idx_cmp_norm_.weight + ape = layer_weight.idx_cmp_ape_.weight + head_dim = self.index_head_dim + else: + norm_weight = layer_weight.compressor_norm_.weight + ape = layer_weight.compressor_ape_.weight + head_dim = self.head_dim return fused_compress_op( kv_score=metadata.kv_score, metadata=metadata, - norm_weight=layer_weight.compressor_norm_.weight, - ape=layer_weight.compressor_ape_.weight, + norm_weight=norm_weight, + ape=ape, eps=self.eps, - head_dim=self.head_dim, + head_dim=head_dim, qk_rope_head_dim=self.qk_rope_head_dim, compress_ratio=self.compress_ratio, cos_table=cos_table, sin_table=sin_table, + output_bf16=self.is_in_indexer, + ) + + +FLASHMLA_INDEX_ALIGN = 64 + + +def _pad_last_dim(x: torch.Tensor, multiple: int = FLASHMLA_INDEX_ALIGN, value: int = -1) -> torch.Tensor: + pad = (-x.shape[-1]) % multiple + if pad == 0: + return x.contiguous() + out = torch.full((*x.shape[:-1], x.shape[-1] + pad), value, dtype=x.dtype, device=x.device) + out[..., : x.shape[-1]] = x + return out.contiguous() + + +class DeepseekV4IndexInfer: + """Model-side builder for the FlashMLA sparse-index metadata. Mirrors deepseek3_2's NsaInfer + *boundary* (the model owns ALL index construction; the attention backend only forwards final + tensors to flash_mla.flash_mla_with_kvcache) but NOT its implementation -- the two share ~no + concrete operators (ds3_2: fp8_mqa_logits over the full ragged kv; dsv4: bf16 einsum over the + compressed c4 entries), hence no inheritance. Owns swa/c128 slot bookkeeping AND the c4 + Lightning-Indexer scoring. Holds only static per-layer config; all per-request data flows in via + args. Invoke from _context/_token_attention_kernel (after compressor.fused_compress, before + *_att) so the c4 einsum/topk/all_reduce keep the same cuda-graph capture position they had when + this lived in the backend.""" + + def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): + self.layer_idx_ = layer_idx + self.compress_ratio = network_config["compress_ratios"][layer_idx] + self.index_topk = network_config["index_topk"] + self.index_head_dim = network_config["index_head_dim"] + self.qk_rope_head_dim = network_config["qk_rope_head_dim"] + self.index_n_heads = network_config["index_n_heads"] + self.tp_world_size_ = tp_world_size + self.tp_index_n_heads = self.index_n_heads // tp_world_size + self.indexer_score_scale = self.index_head_dim ** -0.5 + self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 + # c4 layers own a second compressor (is_in_indexer) that writes the Lightning-Indexer key + # pool every step; the scorer in _c4_indices reads it back via gather_indexer_k. + self.indexer_compressor = ( + CompressorInfer(layer_idx, network_config, tp_world_size, is_in_indexer=True) + if self.compress_ratio == 4 + else None + ) + + def write_indexer_k(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight, cos_table, sin_table): + """c4-only: compress this step's tokens into per-c4-entry indexer keys and pack them into + c4_indexer_pool. MUST run before build_metadata so the scorer's gather_indexer_k reads the + finished entries; runs every step (incl. in the decode graph) so keys accumulate for later + long-context scoring. No-op on c128 / dense layers.""" + if self.compress_ratio != 4: + return + self.indexer_compressor.prepare_states(x, infer_state, layer_weight) + self.indexer_compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) + scratch = self.indexer_compressor._metadata.out_buffer # [T, index_head_dim] bf16 (group-end rows valid) + mem_manager = infer_state.mem_manager + positions = infer_state.position_ids + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + # only group-end tokens finish a c4 entry; mask the rest to -1 so the packer skips them + # (mid-group tokens share the group's c4 slot -> avoids racing a finished slot). + completed = ((positions + 1) % 4 == 0) & (out_slots >= 0) + masked_slots = torch.where(completed, out_slots, torch.full_like(out_slots, -1)).to(torch.int32) + mem_manager.pack_indexer_k_to_cache(self.layer_idx_, masked_slots, scratch) + + def build_metadata(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + """Return the final flash_mla index tensors for this layer's compress variant. swa indices and + the per-token req_idx are layer-independent and precomputed once in init_some_extra_state + (read here); only the c4 scorer / c128 gather is per-layer. The backend pairs these with the + (data-independent, layer-keyed) fp8 cache-byte views it owns.""" + swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) + swa_lengths = infer_state.dsv4_swa_lengths + req_idx = infer_state.dsv4_sparse_req_idx + positions = infer_state.position_ids + extra_indices = extra_lengths = None + if self.compress_ratio == 4: + idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q, idx_weight, req_idx, positions) + elif self.compress_ratio == 128: + extra_indices, extra_lengths = self._c128_indices(infer_state, req_idx, positions) + return { + "swa_indices": swa_indices, + "swa_lengths": swa_lengths, + "extra_indices": extra_indices, + "extra_lengths": extra_lengths, + } + + @staticmethod + def _gather_compress_slots(infer_state, mapping, req_idx, valid, offsets, ratio): + """条目 g 的压缩槽 = full_to_c*[req_to_token[req, (g+1)*ratio-1]](组末 token 的 full 槽位)。 + 无效条目(超出因果长度/HOLD 行)用位置 0 安全 gather 后由调用方按 valid 掩掉。""" + end_pos = offsets[None, :] * ratio + (ratio - 1) + safe_pos = torch.where(valid, end_pos, torch.zeros_like(end_pos)) + full_slots = infer_state.req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_pos] + return mapping[full_slots.long()].to(torch.int32) + + def _c128_indices(self, infer_state: DeepseekV4InferStateInfo, req_idx, positions): + raw_lengths = (positions + 1) // 128 + lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) + max_len = max(1, int(infer_state.max_kv_seq_len) // 128) + offsets = torch.arange(max_len, dtype=torch.long, device=positions.device) + valid = offsets[None, :] < raw_lengths[:, None] + slots = self._gather_compress_slots( + infer_state, infer_state.mem_manager.full_to_c128_indexs, req_idx, valid, offsets, 128 + ) + indices = torch.where(valid, slots, torch.full_like(slots, -1)) + return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() + + def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + cos_tok = infer_state.position_cos_compress + sin_tok = infer_state.position_sin_compress + idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_n_heads, self.index_head_dim) + rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) + idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale + return idx_q, idx_weight + + def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q, idx_weight, req_idx, positions): + """c4(CSA) extra indices: causal all-entries when the entry space fits index_topk, otherwise + Lightning-Indexer scored top-k. Pure tensor ops (decode runs inside cuda graphs).""" + mem_manager = infer_state.mem_manager + raw_lengths = (positions + 1) // 4 + max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) + index_topk = self.index_topk + offsets = torch.arange(max_entries, dtype=torch.long, device=positions.device) + valid = offsets[None, :] < raw_lengths[:, None] + slots = self._gather_compress_slots(infer_state, mem_manager.full_to_c4_indexs, req_idx, valid, offsets, 4) + + if max_entries <= index_topk: + lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) + indices = torch.where(valid, slots, torch.full_like(slots, -1)) + return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() + + score_scale = float(self.indexer_score_scale) + hold_slot = mem_manager.c4_indexer_pool.HOLD_TOKEN_MEMINDEX + safe_slots = torch.where(valid, slots.long(), torch.full_like(slots.long(), hold_slot)) + k = mem_manager.gather_indexer_k(self.layer_idx_, safe_slots.reshape(-1)).view( + positions.shape[0], max_entries, -1 ) + num_tokens, num_heads = idx_q.shape[0], idx_q.shape[1] + score_chunks = [] + chunk = max(1, min(num_tokens, (16 * 1024 * 1024) // max(1, num_heads * max_entries))) + for start in range(0, num_tokens, chunk): + end = min(num_tokens, start + chunk) + scores = torch.einsum("thd,tnd->thn", idx_q[start:end].float(), k[start:end].float()) + scores = F.relu(scores) * score_scale + score_chunks.append((scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1)) + index_scores = torch.cat(score_chunks, dim=0) + if self.tp_world_size_ > 1: + all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) + index_scores = index_scores.masked_fill(~valid, float("-inf")) + top = index_scores.topk(index_topk, dim=-1).indices + top_valid = torch.gather(valid, 1, top) + top_slots = torch.gather(slots.long(), 1, top).to(torch.int32) + indices = torch.where(top_valid, top_slots, torch.full_like(top_slots, -1)) + lengths = torch.clamp(torch.minimum(raw_lengths, torch.full_like(raw_lengths, index_topk)), min=1) + return _pad_last_dim(indices).unsqueeze(1), lengths.to(torch.int32).contiguous() diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py new file mode 100644 index 0000000000..e5b7b80bb4 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py @@ -0,0 +1,75 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _build_swa_index_kernel( + req_idx_ptr, + pos_ptr, + req_to_token_ptr, + req_to_token_stride0, + full_to_swa_ptr, + swa_index_ptr, + swa_length_ptr, + WINDOW: tl.constexpr, + BLOCK_W: tl.constexpr, +): + token_idx = tl.program_id(0) + req = tl.load(req_idx_ptr + token_idx).to(tl.int64) + pos = tl.load(pos_ptr + token_idx).to(tl.int64) + + w = tl.arange(0, BLOCK_W) + w_mask = w < WINDOW + # most-recent-first window, identical to the eager _swa_indices (offset = position - arange). + offset = pos - w + valid = (offset >= 0) & w_mask + safe_offset = tl.where(valid, offset, 0) + full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_offset, mask=valid, other=0).to(tl.int64) + swa_slot = tl.load(full_to_swa_ptr + full_slot, mask=valid, other=-1) + out = tl.where(valid, swa_slot, -1).to(tl.int32) + tl.store(swa_index_ptr + token_idx * WINDOW + w, out, mask=w_mask) + + length = tl.minimum(tl.maximum(pos + 1, 1), WINDOW).to(tl.int32) + tl.store(swa_length_ptr + token_idx, length) + + +def build_swa_index( + req_idx: torch.Tensor, + positions: torch.Tensor, + req_to_token_indexs: torch.Tensor, + full_to_swa_indexs: torch.Tensor, + window: int, +): + """Per-token sliding-window FlashMLA index table, built ONCE per forward (layer-independent: + full_to_swa is a single global map and the window is a model constant, so every layer's swa + indices are identical). Replaces DeepseekV4IndexInfer._swa_indices: for token t at + (req_idx, position) gather the last `window` tokens' full slots via req_to_token, then map + full -> swa; out-of-range positions store -1. + + Returns (swa_index [T, window] int32, swa_length [T] int32). `window` is 128 (a multiple of the + FlashMLA 64 alignment) so no extra pad is needed; the reader adds the s_q axis via unsqueeze(1). + Const output shape (no max_kv_seq_len dependence) makes this cuda-graph-safe to stage from + init_some_extra_state via copy_for_cuda_graph. + """ + # window must stay 64-aligned: the output is the FlashMLA `indices` tensor directly (no separate + # _pad_last_dim), and the extra-cache fork requires the topk dim to be a multiple of 64. + assert window % 64 == 0, f"DeepSeek-V4 sliding_window must be a multiple of 64 for FlashMLA, got {window}" + T = positions.shape[0] + swa_index = torch.empty((T, window), dtype=torch.int32, device=positions.device) + swa_length = torch.empty((T,), dtype=torch.int32, device=positions.device) + if T == 0: + return swa_index, swa_length + _build_swa_index_kernel[(T,)]( + req_idx, + positions, + req_to_token_indexs, + req_to_token_indexs.stride(0), + full_to_swa_indexs, + swa_index, + swa_length, + WINDOW=window, + BLOCK_W=triton.next_power_of_2(window), + num_warps=4, + ) + return swa_index, swa_length From d4dcd8a8924e7218025763447db3abad72d9c520 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 15 Jun 2026 13:00:25 +0000 Subject: [PATCH 020/214] opt --- lightllm/common/basemodel/basemodel.py | 53 +---- .../deepseek4_mem_manager.py | 10 +- lightllm/common/req_manager.py | 40 ++-- .../layer_infer/transformer_layer_infer.py | 192 ++++++++++-------- .../layer_weights/transformer_layer_weight.py | 10 +- .../build_compress_index_dsv4.py | 78 +++++++ .../triton_kernel/gather_c4_indexer_k_dsv4.py | 106 ++++++++++ 7 files changed, 338 insertions(+), 151 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 986802e760..30248d6a21 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -297,6 +297,11 @@ def _init_custom(self): @torch.no_grad() def forward(self, model_input: ModelInput): + # decode 槽位 prep: 放在 to_cuda 前 (b_req_idx/b_seq_len/mem_indexes_cpu 还是原生 CPU 张量), + # 且此刻已在 forward 的 CUDA stream 上 -> 与后续 attention 同流, 无跨流竞态、无 D2H。 + # mem_indexes_cpu is None 时跳过: cudagraph warmup 的输入全在 CUDA 且 b_req_idx 全为 HOLD, prep 本就是 no-op。 + if not model_input.is_prefill and model_input.mem_indexes_cpu is not None: + self.req_manager.prepare_decode(model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes_cpu) model_input.to_cuda() assert model_input.mem_indexes.is_cuda @@ -579,14 +584,6 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) - if hasattr(self.req_manager, "prepare_decode_swa"): - self.req_manager.prepare_decode_swa( - model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes - ) - if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots( - model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes - ) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -608,14 +605,6 @@ def _decode( model_input = self._create_padded_decode_model_input( model_input=model_input, new_batch_size=infer_batch_size ) - if hasattr(self.req_manager, "prepare_decode_swa"): - self.req_manager.prepare_decode_swa( - model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes - ) - if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots( - model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes - ) infer_state = self._create_inferstate(model_input) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -845,6 +834,10 @@ 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): + # decode 槽位 prep: 在 to_cuda 前 (原生 CPU 张量)、且已在 forward 的 CUDA stream 上 (见 forward 注释)。 + for mi in (model_input0, model_input1): + if mi.mem_indexes_cpu is not None: + self.req_manager.prepare_decode(mi.b_req_idx, mi.b_seq_len, mi.mem_indexes_cpu) model_input0.to_cuda() model_input1.to_cuda() assert self.args.enable_tpsp_mix_mode @@ -876,20 +869,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode # 一致,需要按照较高 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) - if hasattr(self.req_manager, "prepare_decode_swa"): - self.req_manager.prepare_decode_swa( - padded_model_input0.b_req_idx, padded_model_input0.b_seq_len, padded_model_input0.mem_indexes - ) - self.req_manager.prepare_decode_swa( - padded_model_input1.b_req_idx, padded_model_input1.b_seq_len, padded_model_input1.mem_indexes - ) - if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots( - padded_model_input0.b_req_idx, padded_model_input0.b_seq_len, padded_model_input0.mem_indexes - ) - self.req_manager.prepare_decode_compress_slots( - padded_model_input1.b_req_idx, padded_model_input1.b_seq_len, padded_model_input1.mem_indexes - ) infer_state0 = self._create_inferstate(padded_model_input0, 0) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, @@ -931,20 +910,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode 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) - if hasattr(self.req_manager, "prepare_decode_swa"): - self.req_manager.prepare_decode_swa( - model_input0.b_req_idx, model_input0.b_seq_len, model_input0.mem_indexes - ) - self.req_manager.prepare_decode_swa( - model_input1.b_req_idx, model_input1.b_seq_len, model_input1.mem_indexes - ) - if hasattr(self.req_manager, "prepare_decode_compress_slots"): - self.req_manager.prepare_decode_compress_slots( - model_input0.b_req_idx, model_input0.b_seq_len, model_input0.mem_indexes - ) - self.req_manager.prepare_decode_compress_slots( - model_input1.b_req_idx, model_input1.b_seq_len, model_input1.mem_indexes - ) infer_state0 = self._create_inferstate(model_input0, 0) copy_kv_index_to_req( self.req_manager.req_to_token_indexs, diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 8d172ec758..32561aa6f7 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -511,8 +511,8 @@ def alloc_swa_prefill( def alloc_swa_decode( self, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, mem_indexes: torch.Tensor, req_to_token_indexs: torch.Tensor, ) -> None: @@ -523,8 +523,8 @@ def alloc_swa_decode( (DSV4 启动参数已拒绝 MTP;支持需按步内顺序分段派生)。""" page = DSV4_SWA_PAGE_SIZE hold_req_id = self.max_request_num - req_list = b_req_idx.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() + req_list = b_req_idx_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() cont_rows, cont_prev_pos, new_rows = [], [], [] for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): req_idx, seq_len = int(req_idx), int(seq_len) @@ -537,7 +537,7 @@ def alloc_swa_decode( cont_prev_pos.append(seq_len - 2) mem_indexes = mem_indexes.cuda().long().reshape(-1) if cont_rows: - req_rows = b_req_idx[cont_rows].long() + req_rows = torch.tensor([req_list[i] for i in cont_rows], dtype=torch.long, device="cuda") prev_full = req_to_token_indexs[req_rows, torch.tensor(cont_prev_pos, device="cuda")].long() prev_slots = self.full_to_swa_indexs[prev_full] # 续槽不变式哨兵: 上一位置必驻留(retain 覆盖)。prep 阶段本就有同步,代价可忽略。 diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index a0faccb00c..9ac1babe23 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -156,6 +156,11 @@ def free_all(self): self.req_list = _ReqLinkedList(self.max_request_num) return + def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): + """每个 decode step 在 to_cuda 之前调用的钩子 (数据为原生 CPU 张量, 且已在 forward 的 + CUDA stream 上)。基类 no-op; 需要 per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" + return + class ReqSamplingParamsManager: """ @@ -479,26 +484,32 @@ def prepare_prefill_swa( self.mem_manager.alloc_swa_prefill(b_req_idx, b_ready_cache_len, b_seq_len, self.req_to_token_indexs) return + def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): + """decode 每步槽位 prep: 先 swa 再 compress。由 BaseModel.forward / microbatch_overlap_decode + 在 to_cuda 之前调用 (CPU 数据 + forward 的 CUDA stream); 不再放在 _decode 里。""" + self.prepare_decode_swa(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) + self.prepare_decode_compress_slots(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) + return + def prepare_decode_swa( self, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, mem_indexes: torch.Tensor, ) -> None: """decode prep: 回收出窗槽并为本步新 token 分配位置对齐的 swa 槽。当前 query 位置 seq_len-1 的窗口是 [seq_len-W, seq_len-1];回收边界额外保留一个 radix 页 - (_swa_retain_len),即位置 < seq_len-retain。先回收再分配。""" + (_swa_retain_len),即位置 < seq_len-retain。先回收再分配。 + seq_len/req_idx 从 CPU 镜像读(host 算术,无 D2H);水位线 _swa_evict_marks 仍是 host 状态。""" assert self.mem_manager is not None if self.sliding_window is not None: retain = self._swa_retain_len() evict_slots = [] - req_list = b_req_idx.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() + req_list = b_req_idx_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() for req_idx, seq_len in zip(req_list, seq_list): - req_idx = int(req_idx) if req_idx == self.HOLD_REQUEST_ID: continue - seq_len = int(seq_len) mark = self._swa_evict_marks[req_idx] if mark < 0: # 未经过 prefill prep 的保守路径: 不回收旧位置,仅推进水位线。 @@ -510,7 +521,7 @@ def prepare_decode_swa( self._swa_evict_marks[req_idx] = evict_end if evict_slots: self.mem_manager.evict_swa(torch.cat(evict_slots)) - self.mem_manager.alloc_swa_decode(b_req_idx, b_seq_len, mem_indexes, self.req_to_token_indexs) + self.mem_manager.alloc_swa_decode(b_req_idx_cpu, b_seq_len_cpu, mem_indexes, self.req_to_token_indexs) return def init_compress_state(self, req_idx: int): @@ -575,23 +586,24 @@ def prepare_prefill_compress_slots( def prepare_decode_compress_slots( self, - b_req_idx: torch.Tensor, - b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, mem_indexes: torch.Tensor, ) -> None: """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 - 组末 full 槽即本步的 mem_index(此刻 req_to_token_indexs 尚未写入本步槽位)。""" + 组末 full 槽即本步的 mem_index(此刻 req_to_token_indexs 尚未写入本步槽位)。 + 从 CPU 镜像读 seq_len/req_idx(host 算术,无 D2H);非关组步 rows 为空 => 不调 _scatter,零同步。""" if self.n_c4 == 0 and self.n_c128 == 0: return - req_list = b_req_idx.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() + req_list = b_req_idx_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() for ratio, n_layers in ((4, self.n_c4), (128, self.n_c128)): if n_layers == 0: continue rows = [ i for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) - if int(req_idx) != self.HOLD_REQUEST_ID and int(seq_len) > 0 and int(seq_len) % ratio == 0 + if req_idx != self.HOLD_REQUEST_ID and seq_len > 0 and seq_len % ratio == 0 ] if rows: self._scatter_compress_slots(ratio, mem_indexes.reshape(-1)[rows]) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index f99df4b00d..ef898c31bc 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -227,12 +227,12 @@ def _context_attention_kernel( self.compressor.prepare_states(x, infer_state, layer_weight) self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) # Write this step's c4 Lightning-Indexer keys (no-op off c4) BEFORE build_metadata so the - # scorer's gather_indexer_k reads fresh+accumulated entries. + # scorer (gather + deep_gemm.fp8_mqa_logits) reads fresh+accumulated entries from the indexer pool. self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) # Build the FINAL flash_mla index tensors here (model side), so att_control is a thin # transport of ready-to-forward tensors -- not indexer raw material. Must stay after # fused_compress (c4 reads the indexer-K pool it writes) and before prefill_att (keeps the - # c4 einsum/topk/all_reduce at the same cuda-graph capture position). + # c4 scorer/topk at the same cuda-graph capture position). meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) att_control = AttControl( nsa_prefill=True, @@ -452,28 +452,18 @@ def fused_compress( ) -FLASHMLA_INDEX_ALIGN = 64 - - -def _pad_last_dim(x: torch.Tensor, multiple: int = FLASHMLA_INDEX_ALIGN, value: int = -1) -> torch.Tensor: - pad = (-x.shape[-1]) % multiple - if pad == 0: - return x.contiguous() - out = torch.full((*x.shape[:-1], x.shape[-1] + pad), value, dtype=x.dtype, device=x.device) - out[..., : x.shape[-1]] = x - return out.contiguous() - - class DeepseekV4IndexInfer: """Model-side builder for the FlashMLA sparse-index metadata. Mirrors deepseek3_2's NsaInfer - *boundary* (the model owns ALL index construction; the attention backend only forwards final - tensors to flash_mla.flash_mla_with_kvcache) but NOT its implementation -- the two share ~no - concrete operators (ds3_2: fp8_mqa_logits over the full ragged kv; dsv4: bf16 einsum over the - compressed c4 entries), hence no inheritance. Owns swa/c128 slot bookkeeping AND the c4 - Lightning-Indexer scoring. Holds only static per-layer config; all per-request data flows in via - args. Invoke from _context/_token_attention_kernel (after compressor.fused_compress, before - *_att) so the c4 einsum/topk/all_reduce keep the same cuda-graph capture position they had when - this lived in the backend.""" + boundary (the model owns ALL index construction; the attention backend only forwards final + tensors to flash_mla.flash_mla_with_kvcache) AND its c4 implementation: hadamard'd fp8 q/K, a + ragged gather of the compressed c4 keys, deep_gemm.fp8_mqa_logits, then topk -- adapted for the + replicated indexer (no gather-q/all_reduce), the c4-compressed entry space, and topk-512 (no + inheritance only because of those data-shape differences). swa metadata is precomputed in + init_some_extra_state; this class owns the c4/c128 entry gather (build_compress_index) AND the c4 + Lightning-Indexer scoring (gather + deep_gemm.fp8_mqa_logits + topk). Holds only static per-layer + config; all per-request data flows in via args. Invoke from _context/_token_attention_kernel + (after compressor.fused_compress, before *_att) so the c4 scorer/topk keep the same cuda-graph + capture position they had when this lived in the backend. The indexer is replicated (no TP collective).""" def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): self.layer_idx_ = layer_idx @@ -483,11 +473,10 @@ def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): self.qk_rope_head_dim = network_config["qk_rope_head_dim"] self.index_n_heads = network_config["index_n_heads"] self.tp_world_size_ = tp_world_size - self.tp_index_n_heads = self.index_n_heads // tp_world_size self.indexer_score_scale = self.index_head_dim ** -0.5 self.indexer_weight_scale = self.indexer_score_scale * self.index_n_heads ** -0.5 # c4 layers own a second compressor (is_in_indexer) that writes the Lightning-Indexer key - # pool every step; the scorer in _c4_indices reads it back via gather_indexer_k. + # pool every step; _c4_indices gathers it back + scores via deep_gemm.fp8_mqa_logits. self.indexer_compressor = ( CompressorInfer(layer_idx, network_config, tp_world_size, is_in_indexer=True) if self.compress_ratio == 4 @@ -496,14 +485,19 @@ def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): def write_indexer_k(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight, cos_table, sin_table): """c4-only: compress this step's tokens into per-c4-entry indexer keys and pack them into - c4_indexer_pool. MUST run before build_metadata so the scorer's gather_indexer_k reads the - finished entries; runs every step (incl. in the decode graph) so keys accumulate for later - long-context scoring. No-op on c128 / dense layers.""" + c4_indexer_pool. MUST run before build_metadata so the scorer (gather + deep_gemm.fp8_mqa_logits) + reads the finished entries; runs every step (incl. in the decode graph) so keys accumulate for + later long-context scoring. No-op on c128 / dense layers.""" if self.compress_ratio != 4: return self.indexer_compressor.prepare_states(x, infer_state, layer_weight) self.indexer_compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) scratch = self.indexer_compressor._metadata.out_buffer # [T, index_head_dim] bf16 (group-end rows valid) + # Rotate K (post norm+rope) by the SAME 1/sqrt(d) Hadamard the q kernel applies, so + # (Hq)·(Hk)=q·k (H orthogonal) and the fp8 quant of K stays accurate. + from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform + + scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5) mem_manager = infer_state.mem_manager positions = infer_state.position_ids out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] @@ -524,8 +518,8 @@ def build_metadata(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer positions = infer_state.position_ids extra_indices = extra_lengths = None if self.compress_ratio == 4: - idx_q, idx_weight = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) - extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q, idx_weight, req_idx, positions) + idx_q_fp8, weights = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: extra_indices, extra_lengths = self._c128_indices(infer_state, req_idx, positions) return { @@ -535,72 +529,98 @@ def build_metadata(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer "extra_lengths": extra_lengths, } - @staticmethod - def _gather_compress_slots(infer_state, mapping, req_idx, valid, offsets, ratio): - """条目 g 的压缩槽 = full_to_c*[req_to_token[req, (g+1)*ratio-1]](组末 token 的 full 槽位)。 - 无效条目(超出因果长度/HOLD 行)用位置 0 安全 gather 后由调用方按 valid 掩掉。""" - end_pos = offsets[None, :] * ratio + (ratio - 1) - safe_pos = torch.where(valid, end_pos, torch.zeros_like(end_pos)) - full_slots = infer_state.req_manager.req_to_token_indexs[req_idx.long()[:, None], safe_pos] - return mapping[full_slots.long()].to(torch.int32) - def _c128_indices(self, infer_state: DeepseekV4InferStateInfo, req_idx, positions): - raw_lengths = (positions + 1) // 128 - lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) - max_len = max(1, int(infer_state.max_kv_seq_len) // 128) - offsets = torch.arange(max_len, dtype=torch.long, device=positions.device) - valid = offsets[None, :] < raw_lengths[:, None] - slots = self._gather_compress_slots( - infer_state, infer_state.mem_manager.full_to_c128_indexs, req_idx, valid, offsets, 128 + from ..triton_kernel.build_compress_index_dsv4 import build_compress_index + + cap = ((max(1, int(infer_state.max_kv_seq_len) // 128) + 63) // 64) * 64 + indices, lengths = build_compress_index( + req_idx, + positions, + infer_state.req_manager.req_to_token_indexs, + infer_state.mem_manager.full_to_c128_indexs, + ratio=128, + cap=cap, ) - indices = torch.where(valid, slots, torch.full_like(slots, -1)) - return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() + return indices.unsqueeze(1), lengths def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + """fp8 indexer q (mirrors deepseek3_2 NsaInfer): wq_b -> rope(last rope dims) -> 1/sqrt(d) + Hadamard -> per-token fp8 quant. Returns (idx_q_fp8 [T,H,d], weights [T,H]); the per-token q + fp8 scale and the head_dim^-0.5 * n_heads^-0.5 score scale are folded into weights -- the + deep_gemm.fp8_mqa_logits contract (fp8 q carries no companion scale). Replicated -> full heads.""" + from lightllm.models.deepseek3_2.triton_kernel.act_quant import act_quant + from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform + cos_tok = infer_state.position_cos_compress sin_tok = infer_state.position_sin_compress - idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.tp_index_n_heads, self.index_head_dim) + idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.index_n_heads, self.index_head_dim) rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) - idx_weight = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale - return idx_q, idx_weight - - def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q, idx_weight, req_idx, positions): - """c4(CSA) extra indices: causal all-entries when the entry space fits index_topk, otherwise - Lightning-Indexer scored top-k. Pure tensor ops (decode runs inside cuda graphs).""" + idx_q = hadamard_transform(idx_q, scale=self.index_head_dim ** -0.5) + idx_q_fp8, q_scale = act_quant(idx_q, self.index_head_dim, None) # fp8 [T,H,d], scale [T,H,1] + weights = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale # [T, H] + weights = weights.unsqueeze(-1) * q_scale # fold per-token q scale + return idx_q_fp8, weights.squeeze(-1).contiguous() + + def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): + """c4 scorer via ds3.2-style gather + deep_gemm.fp8_mqa_logits. Gather each request's causal c4 + keys into a padded-per-request ragged fp8 buffer (k row r*c4_cap+e), score every query token + over its absolute [ks, ke) range, then masked topk-512 -> c4 slots. Fixed shapes (c4_cap pinned + per graph bucket) keep the decode cuda graph capturable.""" mem_manager = infer_state.mem_manager - raw_lengths = (positions + 1) // 4 - max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) index_topk = self.index_topk - offsets = torch.arange(max_entries, dtype=torch.long, device=positions.device) - valid = offsets[None, :] < raw_lengths[:, None] - slots = self._gather_compress_slots(infer_state, mem_manager.full_to_c4_indexs, req_idx, valid, offsets, 4) + max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) + c4_cap = ((max_entries + 63) // 64) * 64 + # entry space fits the budget -> every causal entry is selected; no scoring needed. The + # captured decode graph (graph_max_len -> max_entries > topk) always takes the scorer branch + # below, so this only shortcuts tiny eager contexts. if max_entries <= index_topk: - lengths = torch.clamp(raw_lengths, min=1).to(torch.int32) - indices = torch.where(valid, slots, torch.full_like(slots, -1)) - return _pad_last_dim(indices).unsqueeze(1), lengths.contiguous() - - score_scale = float(self.indexer_score_scale) - hold_slot = mem_manager.c4_indexer_pool.HOLD_TOKEN_MEMINDEX - safe_slots = torch.where(valid, slots.long(), torch.full_like(slots.long(), hold_slot)) - k = mem_manager.gather_indexer_k(self.layer_idx_, safe_slots.reshape(-1)).view( - positions.shape[0], max_entries, -1 + from ..triton_kernel.build_compress_index_dsv4 import build_compress_index + + slots, lengths = build_compress_index( + infer_state.dsv4_sparse_req_idx, + positions, + infer_state.req_manager.req_to_token_indexs, + mem_manager.full_to_c4_indexs, + ratio=4, + cap=c4_cap, + ) + return slots.unsqueeze(1), lengths + + import deep_gemm + from ..triton_kernel.gather_c4_indexer_k_dsv4 import gather_c4_indexer_k_ragged + + b_req_idx = infer_state.b_req_idx + batch = b_req_idx.shape[0] + device = positions.device + c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req + k_fp8, k_scale, ragged_slots = gather_c4_indexer_k_ragged( + mem_manager, + self.layer_idx_, + b_req_idx, + c4_len, + c4_cap, + infer_state.req_manager.req_to_token_indexs, ) - num_tokens, num_heads = idx_q.shape[0], idx_q.shape[1] - score_chunks = [] - chunk = max(1, min(num_tokens, (16 * 1024 * 1024) // max(1, num_heads * max_entries))) - for start in range(0, num_tokens, chunk): - end = min(num_tokens, start + chunk) - scores = torch.einsum("thd,tnd->thn", idx_q[start:end].float(), k[start:end].float()) - scores = F.relu(scores) * score_scale - score_chunks.append((scores * idx_weight[start:end].unsqueeze(-1)).sum(dim=1)) - index_scores = torch.cat(score_chunks, dim=0) - if self.tp_world_size_ > 1: - all_reduce(index_scores, op=dist.ReduceOp.SUM, group=infer_state.dist_group, async_op=False) - index_scores = index_scores.masked_fill(~valid, float("-inf")) - top = index_scores.topk(index_topk, dim=-1).indices - top_valid = torch.gather(valid, 1, top) - top_slots = torch.gather(slots.long(), 1, top).to(torch.int32) - indices = torch.where(top_valid, top_slots, torch.full_like(top_slots, -1)) - lengths = torch.clamp(torch.minimum(raw_lengths, torch.full_like(raw_lengths, index_topk)), min=1) - return _pad_last_dim(indices).unsqueeze(1), lengths.to(torch.int32).contiguous() + # batch position of each query token -> absolute [ks, ke) into the padded buffer. + if infer_state.is_prefill: + token_batch_pos = torch.repeat_interleave( + torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len + ) + else: + token_batch_pos = torch.arange(batch, device=device, dtype=torch.int32) + valid_len = ((positions + 1) // 4).to(torch.int32) # causal candidate count per query + ks = token_batch_pos * c4_cap + ke = ks + valid_len + logits = deep_gemm.fp8_mqa_logits( + idx_q_fp8, (k_fp8, k_scale), weights, ks, ke, clean_logits=False, max_seqlen_k=c4_cap + ) # [T, c4_cap] f32, left-aligned: logits[t, j] = q_t . k[ks[t]+j] + col = torch.arange(c4_cap, device=device) + logits = logits.masked_fill(col.unsqueeze(0) >= valid_len.unsqueeze(1), float("-inf")) + top = logits.topk(index_topk, dim=-1).indices.to(torch.int32) # relative positions in [0, valid_len) + abs_idx = top + ks.unsqueeze(1) # absolute compact row + top_slots = ragged_slots[abs_idx.long()] # compact row -> c4 pool slot + invalid = top >= valid_len.unsqueeze(1) # topk over -inf padding when valid_len < index_topk + top_slots = torch.where(invalid, torch.full_like(top_slots, -1), top_slots) + topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) + return top_slots.unsqueeze(1), topk_lengths.contiguous() diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index d42b20f6e5..581b2b1d96 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -47,7 +47,6 @@ def _parse_config(self): self.is_hash = self.layer_num_ < self.num_hash_layers assert self.n_heads % self.tp_world_size_ == 0 assert self.o_groups % self.tp_world_size_ == 0 - assert self.index_n_heads % self.tp_world_size_ == 0 self.prefix = f"layers.{self.layer_num_}" def _init_weight(self): @@ -139,13 +138,18 @@ def _init_compressor(self): def _init_indexer(self): p = f"{self.prefix}.attn.indexer" - # wq_b is FP8 in the checkpoint -> de-quantized to bf16 at load; column-parallel over index heads. + # The Lightning-Indexer is REPLICATED across TP ranks (like sglang/vllm), not head-sharded: + # q_lora and the attn input are already full on every rank, so each rank scores all + # index_n_heads locally and the c4 top-k is identical everywhere -- no gather/all_reduce. + # wq_b is FP8 in the checkpoint -> de-quantized to bf16 at load. self.idx_wq_b_ = ROWMMWeight( in_dim=self.q_lora_rank, out_dims=[self.index_n_heads * self.index_head_dim], weight_names=f"{p}.wq_b.weight", data_type=self.data_type_, quant_method=self.get_quant_method("idx_wq_b"), + tp_rank=0, + tp_world_size=1, ) self.idx_weights_proj_ = ROWMMWeight( in_dim=self.hidden, @@ -153,6 +157,8 @@ def _init_indexer(self): weight_names=f"{p}.weights_proj.weight", data_type=self.data_type_, quant_method=None, + tp_rank=0, + tp_world_size=1, ) coff = 2 # indexer compressor always uses ratio 4 (overlap) self.idx_cmp_wkv_ = ROWMMWeight( diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py new file mode 100644 index 0000000000..b09192498d --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py @@ -0,0 +1,78 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _build_compress_index_kernel( + req_idx_ptr, + pos_ptr, + req_to_token_ptr, + req_to_token_stride0, + full_to_c_ptr, + index_ptr, + length_ptr, + cap, + RATIO: tl.constexpr, + BLOCK_E: tl.constexpr, +): + t = tl.program_id(0) + eb = tl.program_id(1) + req = tl.load(req_idx_ptr + t).to(tl.int64) + pos = tl.load(pos_ptr + t).to(tl.int64) + raw_len = (pos + 1) // RATIO + + e = eb * BLOCK_E + tl.arange(0, BLOCK_E) + e_mask = e < cap + valid = (e < raw_len) & e_mask + # group-end token of compressed entry e: position e*RATIO + (RATIO-1). + end_pos = e * RATIO + (RATIO - 1) + safe_pos = tl.where(valid, end_pos, 0) + full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_pos, mask=valid, other=0).to(tl.int64) + c_slot = tl.load(full_to_c_ptr + full_slot, mask=valid, other=-1).to(tl.int32) + tl.store(index_ptr + t * cap + e, c_slot, mask=e_mask) + + if eb == 0: + tl.store(length_ptr + t, tl.maximum(raw_len, 1).to(tl.int32)) + + +def build_compress_index( + req_idx: torch.Tensor, + positions: torch.Tensor, + req_to_token_indexs: torch.Tensor, + full_to_c_indexs: torch.Tensor, + ratio: int, + cap: int, +): + """Fused two-level group-end gather for the c4/c128 compressed-entry index tables. + + For token t (at request `req_idx[t]`, absolute `positions[t]`) and compressed entry e: + slot[t, e] = full_to_c[ req_to_token[req, e*ratio + (ratio-1)] ] (the group-end token's full slot) + with slot = -1 where e >= (pos+1)//ratio (beyond the causal compressed length) or where the + full->c map is unset. Returns (index [T, cap] int32, length [T] int32 = clamp((pos+1)//ratio, 1)). + + Replaces the eager _gather_compress_slots/_c128/c4-causal torch chain. `cap` must be a multiple of + 64 (FlashMLA topk alignment); the tiled grid (T, ceil(cap/BLOCK_E)) scales to 1M-context caps. + cuda-graph-safe: cap is fixed per graph bucket, shapes static. + """ + T = positions.shape[0] + index = torch.empty((T, cap), dtype=torch.int32, device=positions.device) + length = torch.empty((T,), dtype=torch.int32, device=positions.device) + if T == 0: + return index, length + BLOCK_E = 256 + grid = (T, triton.cdiv(cap, BLOCK_E)) + _build_compress_index_kernel[grid]( + req_idx, + positions, + req_to_token_indexs, + req_to_token_indexs.stride(0), + full_to_c_indexs, + index, + length, + cap, + RATIO=ratio, + BLOCK_E=BLOCK_E, + num_warps=4, + ) + return index, length diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py new file mode 100644 index 0000000000..6adac4fdd0 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -0,0 +1,106 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gather_c4_indexer_k_kernel( + req_idx_ptr, # [batch] int — req_manager slot per batch position + c4_len_ptr, # [batch] int — number of causal c4 entries per request (= seq_len // ratio) + req_to_token_ptr, + req_to_token_stride0, + full_to_c4_ptr, + SlabFp8_ptr, # c4 indexer pool, viewed as fp8 (flat) + SlabF32_ptr, # same pool, viewed as f32 (flat) + Kout_fp8_ptr, # [batch*c4_cap, HEAD_DIM] fp8 + Kout_scale_ptr, # [batch*c4_cap] f32 + Slots_out_ptr, # [batch*c4_cap] int32 (compact->c4-slot map; -1 for padding) + c4_cap, + RATIO: tl.constexpr, + HEAD_DIM: tl.constexpr, + PAGE_SIZE: tl.constexpr, + BYTES_PER_PAGE: tl.constexpr, + SCALE_OFFSET: tl.constexpr, # page_size * head_dim (byte offset of the scale tail) +): + r = tl.program_id(0) + e = tl.program_id(1) + out_pos = r * c4_cap + e + c4_len = tl.load(c4_len_ptr + r) + if e >= c4_len: + # padding entry: mark slot invalid; K is never read (ke bounds the scorer range). + tl.store(Slots_out_ptr + out_pos, -1) + return + + # group-end token of compressed entry e lives at position e*RATIO + (RATIO-1). + req = tl.load(req_idx_ptr + r).to(tl.int64) + end_tok = e * RATIO + (RATIO - 1) + full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + end_tok).to(tl.int64) + c4_slot = tl.load(full_to_c4_ptr + full_slot).to(tl.int64) + valid = c4_slot >= 0 + + # inline PackedPagePool byte addressing (matches destindex_copy_indexer_k_dsv4 / gather_indexer_k): + # fp8 K at page*bytes_per_page + tok*head_dim; fp32 scale at (page*bytes_per_page + scale_off)//4 + tok. + page = c4_slot // PAGE_SIZE + tok = c4_slot % PAGE_SIZE + data_base = page * BYTES_PER_PAGE + tok * HEAD_DIM + scale_base = (page * BYTES_PER_PAGE + SCALE_OFFSET) // 4 + tok + + offs_d = tl.arange(0, HEAD_DIM) + k_fp8 = tl.load(SlabFp8_ptr + data_base + offs_d, mask=valid, other=0.0) + k_scale = tl.load(SlabF32_ptr + scale_base, mask=valid, other=0.0) + tl.store(Kout_fp8_ptr + out_pos * HEAD_DIM + offs_d, k_fp8) + tl.store(Kout_scale_ptr + out_pos, k_scale) + tl.store(Slots_out_ptr + out_pos, tl.where(valid, c4_slot, -1).to(tl.int32)) + + +@torch.no_grad() +def gather_c4_indexer_k_ragged( + mem_manager, + layer_index: int, + b_req_idx: torch.Tensor, + c4_len: torch.Tensor, + c4_cap: int, + req_to_token_indexs: torch.Tensor, +): + """Gather each request's causal c4 indexer keys into a padded-per-request ragged buffer for the + deep_gemm fp8_mqa_logits scorer (mirrors deepseek3_2's extract_indexer_ks, but reads our + PackedPagePool by c4 slot instead of a token-indexed [N,1,132] buffer). + + For batch position r and compressed entry e in [0, c4_len[r]): + c4_slot = full_to_c4[req_to_token[b_req_idx[r], e*ratio + (ratio-1)]] + The raw fp8 key + f32 scale at that slot land at row r*c4_cap + e of the output (so query token t + of request r reads keys [r*c4_cap, r*c4_cap + (pos+1)//ratio) -- absolute ks/ke offsets the caller + builds). Returns (k_fp8 [batch*c4_cap, HEAD_DIM] fp8, k_scale [batch*c4_cap] f32, slots + [batch*c4_cap] int32 = compact-row -> c4 pool slot, -1 for padding). Fixed shapes -> cuda-graph + safe (c4_cap is pinned per graph bucket); the padding region is never read by the scorer. + """ + pool = mem_manager.c4_indexer_pool + head_dim = mem_manager.indexer_head_dim + buf = pool.get_layer_buffer(mem_manager.layer_to_c4_idx[layer_index]).view(-1) + slab_fp8 = buf.view(torch.float8_e4m3fn) + slab_f32 = buf.view(torch.float32) + batch = b_req_idx.shape[0] + n = batch * c4_cap + k_fp8 = torch.empty((n, head_dim), dtype=torch.float8_e4m3fn, device=buf.device) + k_scale = torch.empty((n,), dtype=torch.float32, device=buf.device) + slots = torch.empty((n,), dtype=torch.int32, device=buf.device) + _gather_c4_indexer_k_kernel[(batch, c4_cap)]( + b_req_idx, + c4_len, + req_to_token_indexs, + req_to_token_indexs.stride(0), + mem_manager.full_to_c4_indexs, + slab_fp8, + slab_f32, + k_fp8, + k_scale, + slots, + c4_cap, + RATIO=4, + HEAD_DIM=head_dim, + PAGE_SIZE=pool.page_size, + BYTES_PER_PAGE=pool.bytes_per_page, + SCALE_OFFSET=pool.scale_offset_in_page, + num_warps=1, + ) + return k_fp8, k_scale, slots From 62c16d5675787b308395724ca1fea5689665f622 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 15 Jun 2026 13:27:49 +0000 Subject: [PATCH 021/214] opt --- lightllm/common/basemodel/basemodel.py | 8 +----- .../layer_infer/transformer_layer_infer.py | 10 +++---- .../layer_weights/transformer_layer_weight.py | 26 +++++++++---------- lightllm/models/deepseek_v4/model.py | 5 ---- .../triton_kernel/gather_c4_indexer_k_dsv4.py | 11 +++++--- 5 files changed, 24 insertions(+), 36 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 30248d6a21..b2782d3ed2 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -640,13 +640,7 @@ def prefill_func(input_tensors, infer_state): handle_token_num = infer_state.input_ids.shape[0] - can_run_prefill_graph = self.prefill_graph is not None and self.prefill_graph.can_run( - handle_token_num=handle_token_num - ) - if can_run_prefill_graph and hasattr(self, "_can_run_prefill_cudagraph"): - can_run_prefill_graph = self._can_run_prefill_cudagraph(infer_state, handle_token_num) - - if can_run_prefill_graph: + if self.prefill_graph is not None and self.prefill_graph.can_run(handle_token_num=handle_token_num): finded_handle_token_num = self.prefill_graph.find_closest_graph_handle_token_num( handle_token_num=handle_token_num ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index ef898c31bc..6ebc0b2856 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,5 +1,4 @@ import torch -import torch.nn.functional as F import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl @@ -306,14 +305,13 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV if not self.enable_ep_moe: x = self._tpsp_allgather(input=x, infer_state=infer_state) - gw = layer_weight.gate_weight_.mm_param.weight - logits = F.linear(x.float(), gw.float()).contiguous() + logits = layer_weight.gate_weight_.mm(x.float()).contiguous() weights, indices = self._select_experts(logits, infer_state, layer_weight) # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 - g = layer_weight.shared_gate_.mm(x).float().clamp(max=self.swiglu_limit) - u = layer_weight.shared_up_.mm(x).float().clamp(min=-self.swiglu_limit, max=self.swiglu_limit) - shared = layer_weight.shared_down_.mm((F.silu(g) * u).to(x.dtype)) + # 复用 Llama 的 _ffn_tp: fused gate_up matmul + silu_and_mul triton kernel,无 swiglu clamp, + # 对齐参考 DeepseekV4MLP(=LlamaMLP)。swiglu_limit clamp 只属于 routed 专家 (见 _routed_experts)。 + shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) routed = self._routed_experts(x, weights, indices, layer_weight) if self.enable_ep_moe: if self.tp_world_size_ > 1: diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 581b2b1d96..5896027b38 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -189,12 +189,14 @@ def _init_indexer(self): # ------------------------------------------------------------------ moe def _init_moe(self): p = f"{self.prefix}.ffn" - # router gate (replicated) + # router gate (replicated). Stored as fp32: the topk_hash_softplus_sqrt router wants fp32 logits, + # so keep the gate matmul in fp32 — but store the (constant) weight as fp32 once here instead of + # re-casting it to fp32 on every forward in _ffn. self.gate_weight_ = ROWMMWeight( in_dim=self.hidden, out_dims=[self.n_routed_experts], weight_names=f"{p}.gate.weight", - data_type=self.data_type_, + data_type=torch.float32, quant_method=None, tp_rank=0, tp_world_size=1, @@ -209,23 +211,19 @@ def _init_moe(self): self.gate_bias_ = ParameterWeight( weight_name=f"{p}.gate.bias", data_type=torch.float32, weight_shape=(self.n_routed_experts,) ) - # shared expert (dense, bf16 after de-quant): w1=gate, w3=up (row), w2=down (col) + # shared expert (dense, bf16 after de-quant): w1=gate, w3=up fused (row), w2=down (col). + # Named gate_up_proj/down_proj so the inherited Llama `_ffn_tp` (fused gate_up matmul + + # silu_and_mul triton kernel, no swiglu clamp) drives it directly. Order [w1, w3] = [gate, up] + # matches silu_and_mul_fwd's blocked layout (first half gate, second half up). sp = f"{p}.shared_experts" - self.shared_gate_ = ROWMMWeight( + self.gate_up_proj = ROWMMWeight( in_dim=self.hidden, - out_dims=[self.moe_inter], - weight_names=f"{sp}.w1.weight", + out_dims=[self.moe_inter, self.moe_inter], + weight_names=[f"{sp}.w1.weight", f"{sp}.w3.weight"], data_type=self.data_type_, quant_method=self.get_quant_method("shared_gate"), ) - self.shared_up_ = ROWMMWeight( - in_dim=self.hidden, - out_dims=[self.moe_inter], - weight_names=f"{sp}.w3.weight", - data_type=self.data_type_, - quant_method=self.get_quant_method("shared_up"), - ) - self.shared_down_ = COLMMWeight( + self.down_proj = COLMMWeight( in_dim=self.moe_inter, out_dims=[self.hidden], weight_names=f"{sp}.w2.weight", diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 63430e548b..f9d341e395 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -116,11 +116,6 @@ def _init_cudagraph(self): self.graph_max_len_in_batch = DSV4_DECODE_CUDAGRAPH_MAX_LEN return super()._init_cudagraph() - def _can_run_prefill_cudagraph(self, infer_state: DeepseekV4InferStateInfo, handle_token_num): - if infer_state.prefix_total_token_num == 0: - return True - return False - def _init_att_backend(self): args = get_env_start_args() if args.llm_kv_type == "None": diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py index 6adac4fdd0..b6e6ba751a 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -22,9 +22,12 @@ def _gather_c4_indexer_k_kernel( BYTES_PER_PAGE: tl.constexpr, SCALE_OFFSET: tl.constexpr, # page_size * head_dim (byte offset of the scale tail) ): - r = tl.program_id(0) - e = tl.program_id(1) - out_pos = r * c4_cap + e + # entry index on grid-X (limit ~2^31), batch on grid-Y (<= running_max_req_size): c4_cap reaches + # 65536 at 256K context, which would blow the 65535 grid-Y cap if entries were the grid-Y axis. + e = tl.program_id(0) + r = tl.program_id(1) + # int64: out_pos*HEAD_DIM can exceed int32 at high batch + long context (read side is already int64). + out_pos = r.to(tl.int64) * c4_cap + e c4_len = tl.load(c4_len_ptr + r) if e >= c4_len: # padding entry: mark slot invalid; K is never read (ke bounds the scorer range). @@ -84,7 +87,7 @@ def gather_c4_indexer_k_ragged( k_fp8 = torch.empty((n, head_dim), dtype=torch.float8_e4m3fn, device=buf.device) k_scale = torch.empty((n,), dtype=torch.float32, device=buf.device) slots = torch.empty((n,), dtype=torch.int32, device=buf.device) - _gather_c4_indexer_k_kernel[(batch, c4_cap)]( + _gather_c4_indexer_k_kernel[(c4_cap, batch)]( b_req_idx, c4_len, req_to_token_indexs, From 69824d0059bd0dc9c355dbcac9b79ba715d0d932 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 15 Jun 2026 13:34:37 +0000 Subject: [PATCH 022/214] delete launch.sh --- launch.sh | 53 ----------------------------------------------------- 1 file changed, 53 deletions(-) delete mode 100644 launch.sh diff --git a/launch.sh b/launch.sh deleted file mode 100644 index b9c10d3f0a..0000000000 --- a/launch.sh +++ /dev/null @@ -1,53 +0,0 @@ -# DeepSeek-V4-Flash serving (run inside the lightllm container, repo mounted at /data/wanzihao/lightllm-ds4). -# Verified 2026-06-11: smoke + gsm8k pass with this configuration (prompt cache ENABLED, decode -# cudagraph ENABLED; gsm8k 100q/128: cold 0.960/112s, warm 0.970/23.5s with 100% cache hits — -# vs eager cold 0.970/141s, warm 0.960/50s; batch-1 decode 20.4ms/token vs 142ms eager). -# -# Required env/flags and why: -# LOADWORKER=16 - parallel weight loading (~5x faster startup). -# PYTHONPATH sglang - _get_qkv / compressor reuse sglang.jit_kernel.dsv4 (fused_q_norm_rope, compress_old). -# --batch_max_tokens 8192 - FlashMLA get_decoding_sched_meta rejects >8192 rows per call (probed: 8192 OK, 12288 fails). -# kv pool sizing: auto-profiled from mem_fraction. The fp4 marlin MoE weights materialize their -# CUDA marlin-layout buffers at construction (MXFP4MoEQuantizationMethod._create_weight), so the -# profile sees the true weight footprint on any GPU/config. --max_total_token_num overrides. -# decode cudagraph ENABLED - the v5 decode path is graph-safe: slot alloc/scatter in prep (outside -# graph), forward is pure gathers, HOLD padding rows redirect to HOLD slots. CORRECTNESS NOTE: -# FlashMLASchedMeta is lazily planned at first kernel call and written back onto the (shared) -# decode att state; the capture warmup pass would bake a dummy-content plan into the graph -# (gsm8k dropped to 0.74 with coherent-but-runaway generations). reset_sched_meta_for_capture() -# in cuda_graph._capture_decode re-plans INSIDE the captured region so every replay re-plans. -# DSV4 caps graph max_len_in_batch at 8192; longer decode batches fall back to eager. -# --enable_prefill_cudagraph + --prefill_cudagraph_max_handle_token 2048 - graph-sandwich prefill: -# graphs capture only the per-token dense ops; attention/compressor/indexer run eagerly between -# graph segments (att_func), so host-side planning and .tolist() prep never enter capture. Only -# cold prefills (prefix_total_token_num == 0, model gate) of <= 2048 new tokens replay; cache-hit -# and large batched prefills stay eager. Buckets are padded with a HOLD tail request whose -# attention output MUST be zeroed (infer_struct._dsv4_prefill_pad_q_len): pad rows read the -# racing HOLD slot, and nondeterministic pad hiddens perturb real rows via MoE expert batching -# (ulp-level, chaotically amplified ~1.9x/layer to O(1) by layer ~16 -> greedy token flips). -# Residual caveat: padded-vs-unpadded expert-batch composition still shifts reductions by ulps, -# same class as decode bucket padding; run-to-run determinism is anyway bounded by the fp4 -# marlin MoE kernel itself (probabilistic 1-ulp reduction-order noise measured eager-vs-eager). -# Acceptance is therefore statistical (gsm8k parity), not bitwise. -# --disable_flashinfer_allreduce - flashinfer cuda_ipc resolves libcudart to tilelang's stub (undefined cudaDeviceReset); symm-mem allreduce is used instead. -# -# One-time container setup already applied (survives until container rebuild): -# pip install ipython (sglang import dependency) -# site-packages/vllm: layers/mhc.py + kernels/mhc/ + _tilelang_ops.py overlaid from /data/wanzihao/vllm (mhc_pre_tilelang ops; original kept at layers/mhc.py.bak) -# -# original: python -m lightllm.server.api_server --model_dir /data/models/DeepSeek-V4-Flash --tp 4 --enable_prefill_cudagraph - -# repo root = this script's directory, so the same file works in the main tree and in worktrees -# (a hardcoded tree path here once made a worktree launch silently serve main-tree code). -REPO_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" - -LOADWORKER=16 \ -PYTHONPATH="${REPO_DIR}":/data/wanzihao/sglang/python \ -python -m lightllm.server.api_server \ - --model_dir /data/models/DeepSeek-V4-Flash \ - --tp 4 \ - --batch_max_tokens 8192 \ - --disable_flashinfer_allreduce \ - --enable_prefill_cudagraph \ - --prefill_cudagraph_max_handle_token 2048 \ - --port 8000 From df70ecbf6a09bde7be6fccd9e0467794828a9c5b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 15 Jun 2026 14:15:11 +0000 Subject: [PATCH 023/214] fix --- lightllm/models/deepseek_v4/model.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index f9d341e395..6b5d5f9e38 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -148,8 +148,11 @@ def _init_to_get_rotary(self): cfg = self.config rs = cfg.get("rope_scaling", {}) or {} dim = cfg["qk_rope_head_dim"] + # The rope tables MUST span every absolute position any request can produce (the served + # max_req_total_len / max_position_embeddings, up to 1M). Capping them shorter makes + # init_some_extra_state's index_select(cos/sin, position_ids) read OOB past the table at + # contexts beyond the cap (device-side assert / crash). ~268MB total at 1M, fp32x32 x4 views. max_seq = max(int(self.max_seq_length), int(cfg.get("max_position_embeddings", 8192))) - max_seq = min(max_seq, 1 << 18) # cap table size (256K) for correctness-first freq_exponents = torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim positions = torch.arange(max_seq, dtype=torch.float32, device="cuda") From 1ad981d06809631c1cd91fa8780440522e0894d1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 16 Jun 2026 01:01:55 +0000 Subject: [PATCH 024/214] restore --- lightllm/__init__.py | 27 --------------------------- 1 file changed, 27 deletions(-) diff --git a/lightllm/__init__.py b/lightllm/__init__.py index 8e515afb70..e9ba6f3041 100644 --- a/lightllm/__init__.py +++ b/lightllm/__init__.py @@ -1,31 +1,4 @@ from lightllm.utils.device_utils import is_musa - -def _patch_mp_resource_tracker_for_semaphore(): - from multiprocessing import resource_tracker - - if getattr(resource_tracker, "_lightllm_ignore_semaphore", False): - return - - orig_register = resource_tracker.register - orig_unregister = resource_tracker.unregister - - def register(name, rtype): - if rtype == "semaphore": - return - return orig_register(name, rtype) - - def unregister(name, rtype): - if rtype == "semaphore": - return - return orig_unregister(name, rtype) - - resource_tracker.register = register - resource_tracker.unregister = unregister - resource_tracker._lightllm_ignore_semaphore = True - - -_patch_mp_resource_tracker_for_semaphore() - if is_musa(): import torchada # noqa: F401 From 7b17bb554d2ac324a7796132c78b270675398b33 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 16 Jun 2026 01:47:41 +0000 Subject: [PATCH 025/214] support parser --- lightllm/models/deepseek_v4/model.py | 5 ++++ lightllm/server/api_cli.py | 1 + lightllm/server/build_prompt.py | 7 +++++ lightllm/server/function_call_parser.py | 37 +++++++++++++++++++++++-- lightllm/utils/config_utils.py | 8 ++++-- 5 files changed, 53 insertions(+), 5 deletions(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 6b5d5f9e38..887e72f433 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -191,6 +191,11 @@ def _init_to_get_rotary(self): class DeepSeekV4Tokenizer: """Tokenizer wrapper for DeepSeek-V4's Python prompt encoding.""" + # DeepSeek-V4 has a per-request thinking mode (...) toggled via + # chat_template_kwargs={"thinking": true}. It has no Jinja chat_template string, + # so advertise thinking support explicitly for tokenizer_supports_force_thinking(). + supports_thinking = True + def __init__(self, tokenizer, model_dir): self.tokenizer = tokenizer self.model_dir = model_dir diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 70d5c72ac3..7a9ddb0968 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -169,6 +169,7 @@ def make_argument_parser() -> argparse.ArgumentParser: "qwen", "deepseekv31", "deepseekv32", + "deepseekv4", "glm47", "kimi_k2", "qwen3_coder", diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index 96b4d040bb..be6042fa56 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -53,6 +53,13 @@ def tokenizer_supports_force_thinking() -> bool: assert tokenizer is not None + # Tokenizers that encode prompts in Python (e.g. DeepSeek-V4) have no Jinja + # chat_template string to inspect, so advertise thinking support via an + # explicit attribute instead. + if getattr(tokenizer, "supports_thinking", False): + logger.info("tokenizer_supports_force_thinking : True (explicit attribute)") + return True + try: ans = "thinking" in tokenizer.chat_template or "enable_thinking" in tokenizer.chat_template logger.debug(f"chat_template: {tokenizer.chat_template}") diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index f204c154ed..9213f4c7d4 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -40,6 +40,7 @@ "[TOOL_CALLS]", "<|tool▁calls▁begin|>", "<|DSML|function_calls>", + "<|DSML|tool_calls>", ] @@ -1480,11 +1481,14 @@ class DeepSeekV32Detector(BaseFormatDetector): Reference: https://huggingface.co/deepseek-ai/DeepSeek-V3.2 """ - def __init__(self): + def __init__(self, block_name: str = "function_calls"): super().__init__() self.dsml_token = "|DSML|" - self.bot_token = f"<{self.dsml_token}function_calls>" - self.eot_token = f"" + # DeepSeek V3.2 wraps tool calls in a `function_calls` block; V4 uses + # `tool_calls`. Only the outer block name differs — the invoke/parameter + # grammar is identical — so subclasses just override block_name. + self.bot_token = f"<{self.dsml_token}{block_name}>" + self.eot_token = f"" self.invoke_start_prefix = f"<{self.dsml_token}invoke" self.invoke_end_token = f"" self.param_end_token = f"" @@ -1962,6 +1966,32 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami self._buffer = current_text[eot_pos + len(self.eot_token) :].lstrip() +class DeepSeekV4Detector(DeepSeekV32Detector): + """ + Detector for DeepSeek V4 model function call format using DSML. + + Identical grammar to V3.2 (``<|DSML|invoke name="...">`` blocks with + ``<|DSML|parameter name="k" string="true|false">v`` + tags), except the outer block is named ``tool_calls`` instead of + ``function_calls`` — matching the model's own encoding (encoding_dsv4.py: + ``tool_calls_block_name = "tool_calls"``) and system prompt. + + Format Structure: + ``` + <|DSML|tool_calls> + <|DSML|invoke name="get_weather"> + <|DSML|parameter name="location" string="true">Hangzhou + + + ``` + + Reference: https://huggingface.co/deepseek-ai/DeepSeek-V4 + """ + + def __init__(self): + super().__init__(block_name="tool_calls") + + class FunctionCallParser: """ Parser for function/tool calls in model outputs. @@ -1975,6 +2005,7 @@ class FunctionCallParser: "deepseekv3": DeepSeekV3Detector, "deepseekv31": DeepSeekV31Detector, "deepseekv32": DeepSeekV32Detector, + "deepseekv4": DeepSeekV4Detector, "glm47": Glm47Detector, "kimi_k2": KimiK2Detector, "llama3": Llama32Detector, diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 21df2130e0..dcaf7315dd 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -444,6 +444,10 @@ def get_tool_call_parser_for_model(model_path: str) -> Optional[str]: if model_type == "deepseek_v32": return "deepseekv32" + # DeepSeek V4 + if model_type == "deepseek_v4": + return "deepseekv4" + return None @@ -468,8 +472,8 @@ def get_reasoning_parser_for_model(model_path: str) -> Optional[str]: ]: return "qwen3" - # DeepSeek V3 - if model_type in ["deepseek_v3", "deepseek_v31", "deepseek_v32"]: + # DeepSeek V3 / V4 (share the ... reasoning format, request-gated) + if model_type in ["deepseek_v3", "deepseek_v31", "deepseek_v32", "deepseek_v4"]: return "deepseek-v3" # DeepSeek R1 From 6837abd33ff5f3ecf105a465320fb0fb7ff48f9f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 16 Jun 2026 04:05:42 +0000 Subject: [PATCH 026/214] fix --- .../basemodel/attention/nsa/fp8_flashmla_sparse.py | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 03595c1bce..dc18ecf4ba 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -25,18 +25,6 @@ def _pad_q_heads(q_4d: torch.Tensor, attn_sink: torch.Tensor): return q_pad, sink_pad, h_q -class DeepseekV4MissingOperatorError(RuntimeError): - pass - - -def _missing_attention_op(feature: str) -> None: - raise DeepseekV4MissingOperatorError( - f"DeepSeek-V4 {feature} has no production batch operator. The flashmla_kvcache path " - f"(packed swa/c4/c128 pools + paged compressor + indexer top-k) is the supported route; " - f"this legacy/non-flashmla entry point was never wired and is fenced on purpose." - ) - - def _view_dsv4_flashmla_cache(layer_buffer: torch.Tensor, page_size: int) -> torch.Tensor: from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_MLA_BYTES_PER_TOKEN From 02a24ce306bfa31350dd2117f66715e67e8dbde9 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 18 Jun 2026 01:49:41 +0000 Subject: [PATCH 027/214] add c4 paged indexes --- .../deepseek4_mem_manager.py | 55 ++++-- lightllm/common/req_manager.py | 166 ++++++++++++++++-- .../layer_infer/transformer_layer_infer.py | 89 +++++++++- .../triton_kernel/gather_c4_indexer_k_dsv4.py | 91 ++++++++++ .../server/router/model_infer/infer_batch.py | 3 +- 5 files changed, 373 insertions(+), 31 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 32561aa6f7..cfb149dcec 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -31,6 +31,7 @@ DSV4_SWA_PAGE_SIZE = 128 DSV4_C4_PAGE_SIZE = 64 DSV4_C128_PAGE_SIZE = 2 +DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 # compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 离线聚合为每页 1 组。 DSV4_C4_STATE_RING = 8 DSV4_C128_STATE_RING = 128 @@ -194,11 +195,12 @@ def __init__( # ------------------------------------------------------------------ sizing def _swa_per_req_budget(self) -> int: # 活跃请求保留 window + 一个 radix 页(req_manager._swa_retain_len: 让最近完成的 - # 128 边界的结尾页恒驻留,prompt cache 插入门才能放行),即 v5 §2 的「活跃窗口跨页 ≤2」。 - return int(self.sliding_window) + DSV4_SWA_PAGE_SIZE + # prompt-cache 边界的结尾页恒驻留)。V4 的 prompt-cache 边界取 256 token, + # 避免 radix 共享前缀落在 c4 物理页(64 c4 entry = 256 token)中间。 + return int(self.sliding_window) + DSV4_PROMPT_CACHE_PAGE_SIZE def _planned_swa_size(self, full_size: int) -> int: - # swa 池按页分配(页 = 128 = sliding_window = radix 页),容量向上取整到整页。 + # swa 池按页分配(页 = 128 = sliding_window),容量向上取整到整页。 if self.max_request_num is None or self.sliding_window is None: return _ceil_div(full_size, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE cap = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num @@ -335,6 +337,8 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.c4_pool: Optional[PackedPagePool] = None self.c4_indexer_pool: Optional[PackedPagePool] = None self.c4_allocator: Optional[KvCacheAllocator] = None + self.c4_page_allocator: Optional[KvCacheAllocator] = None + self.c4_page_live_count: Optional[torch.Tensor] = None self.c128_pool: Optional[PackedPagePool] = None self.c128_allocator: Optional[KvCacheAllocator] = None self.c4_state_buffer: Optional[torch.Tensor] = None @@ -360,9 +364,12 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): data_bytes=self.indexer_head_dim, scale_bytes=DSV4_INDEXER_SCALE_BYTES, ) - self.c4_allocator = KvCacheAllocator( - self.c4_size, shared_name=f"{server}_dsv4_c4_can_use_token_num_{rank_in_node}" + self.c4_num_pages = self.c4_size // DSV4_C4_PAGE_SIZE + assert self.c4_num_pages > 0, "DeepSeek-V4 c4 pool must have at least one usable full page" + self.c4_page_allocator = KvCacheAllocator( + self.c4_num_pages, shared_name=f"{server}_dsv4_c4_can_use_page_num_{rank_in_node}" ) + self.c4_page_live_count = torch.zeros((self.c4_pool.num_pages,), dtype=torch.int32, device="cuda") self.full_to_c4_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") self.full_to_c4_indexs[size] = self.c4_pool.HOLD_TOKEN_MEMINDEX # c4 compressor 在途状态(attention + indexer): swa 页派生寻址(翻译③),随 swa 页 @@ -588,11 +595,36 @@ def _evict_compress(self, full_slots: torch.Tensor, mapping: torch.Tensor, alloc mapping[full_slots[valid]] = -1 return + def alloc_c4_pages(self, need_pages: int) -> torch.Tensor: + assert self.c4_page_allocator is not None, "DeepSeek-V4 c4 page allocator is not initialized" + return self.c4_page_allocator.alloc(need_pages) + + def count_c4_slots(self, c4_slots: torch.Tensor, delta: int) -> torch.Tensor: + """按 c4 slot 所在页更新存活计数,返回触达的页(去重)。""" + assert self.c4_page_live_count is not None, "DeepSeek-V4 c4 page live count is not initialized" + pages = torch.div(c4_slots.long(), DSV4_C4_PAGE_SIZE, rounding_mode="floor") + ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) + self.c4_page_live_count.index_add_(0, pages, ones) + return torch.unique(pages) + def evict_c4(self, full_slots: torch.Tensor) -> None: """回收 full 槽位(组末 token)映射的 c4 槽。非组末/未映射(-1)的槽位跳过。""" - if self.c4_allocator is None or full_slots.numel() == 0: + if self.c4_page_allocator is None or full_slots.numel() == 0: return - self._evict_compress(full_slots, self.full_to_c4_indexs, self.c4_allocator) + full_slots = full_slots.cuda().long().reshape(-1) + full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) + if full_slots.numel() == 0: + return + slots = self.full_to_c4_indexs[full_slots] + valid = slots >= 0 + valid_slots = slots[valid] + if valid_slots.numel() == 0: + return + self.full_to_c4_indexs[full_slots[valid]] = -1 + touched = self.count_c4_slots(valid_slots, -1) + empty = touched[self.c4_page_live_count[touched] == 0] + if empty.numel() > 0: + self.c4_page_allocator.free(empty.to(torch.int32)) return def evict_c128(self, full_slots: torch.Tensor) -> None: @@ -623,8 +655,9 @@ def free_all(self): self.swa_page_live_count.zero_() self.full_to_swa_indexs.fill_(-1) self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - if self.c4_allocator is not None: - self.c4_allocator.free_all() + if self.c4_page_allocator is not None: + self.c4_page_allocator.free_all() + self.c4_page_live_count.zero_() self.full_to_c4_indexs.fill_(-1) self.full_to_c4_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c4_pool.HOLD_TOKEN_MEMINDEX if self.c128_allocator is not None: @@ -634,13 +667,13 @@ def free_all(self): return def alloc_c4(self, need_size) -> torch.Tensor: - return self.c4_allocator.alloc(need_size) + raise AssertionError("DeepSeek-V4 c4 uses page-safe allocation; call alloc_c4_pages instead") def alloc_c128(self, need_size) -> torch.Tensor: return self.c128_allocator.alloc(need_size) def free_c4(self, free_index) -> None: - self.c4_allocator.free(free_index) + raise AssertionError("DeepSeek-V4 c4 uses page live-count release; call evict_c4 instead") def free_c128(self, free_index) -> None: self.c128_allocator.free(free_index) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 9ac1babe23..5e7c4f96dd 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -17,6 +17,10 @@ from lightllm.common.linear_att_cache_manager.linear_att_buffer_manager import ( LinearAttCacheManager, ) +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_C4_PAGE_SIZE, + DSV4_PROMPT_CACHE_PAGE_SIZE, +) if TYPE_CHECKING: from lightllm.server.router.model_infer.infer_batch import InferReq @@ -30,10 +34,11 @@ class DeepseekV4PromptCachePayload: 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4/c128 compressor 状态以 swa - 页派生寻址(随 swa 页生灭,命中零拷贝续算)。c128 partial state 不跨 radix 的 128 边界保存。 + 页派生寻址(随 swa 页生灭,命中零拷贝续算)。prompt cache 对齐到 256 token, + 避免共享前缀停在 c4 物理页中间。 * ``swa_page_valid``: cpu bool [cache_len // page],插入时按当下 full_to_swa 映射写定 - (页内 128 个映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 128 边界, + (页内 token 映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 page 边界, swa 压力阀回收节点页时清零。""" cache_len: int @@ -61,7 +66,7 @@ def invalidate_swa_pages(self, payload: DeepseekV4PromptCachePayload) -> None: return def valid_match_length(self, payload: Optional[DeepseekV4PromptCachePayload], natural_len: int) -> int: - """radix 匹配裁剪: 返回 <= natural_len 的最大 128 边界 L',使结尾页(bitmap[L'/128-1])有效。 + """radix 匹配裁剪: 返回 <= natural_len 的最大 prompt-cache 边界 L',使结尾页有效。 有效性可能非单调(owner 生前从左驱逐、后续阀从尾回收),按候选边界回查 bitmap; 中段 invalid 页不挡更靠后的有效命中(注意力只回看最后一个窗口)。""" @@ -441,10 +446,9 @@ def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): def _swa_retain_len(self) -> int: """出窗回收的保留长度 = window + 一个 radix 页。 - 多留一页使「最近一个完成的 128 边界」的结尾页恒驻留: prompt cache 只能在 floor(cur/128) - 边界入树(radix page=128),若回收只留 window,则任何非对齐时刻该边界的结尾页都已被 - 部分回收,插入门会把所有插入裁到 0(prompt cache 形同虚设)。预算即 v5 §2 的每请求 - 「活跃窗口跨页 ≤2」。驻留证明要求 window >= page-1(DSV4 实际 window == page == 128)。""" + 多留一页使「最近一个完成的 prompt-cache 边界」的结尾页恒驻留: 若回收只留 window, + 则任何非对齐时刻该边界的结尾页都已被部分回收,插入门会把所有插入裁到 0。 + V4 prompt-cache 页取 256 token,正好覆盖一个 c4 物理页对应的 token 范围。""" return int(self.sliding_window) + self.get_prompt_cache_page_size() def prepare_prefill_swa( @@ -535,11 +539,133 @@ def init_compress_state(self, req_idx: int): def _compress_mapping_alloc(self, ratio: int): assert self.mem_manager is not None, "DeepSeek-V4 mem manager is not bound yet" if ratio == 4: - return self.mem_manager.full_to_c4_indexs, self.mem_manager.alloc_c4 + raise AssertionError("DeepSeek-V4 c4 uses page-safe allocation") if ratio == 128: return self.mem_manager.full_to_c128_indexs, self.mem_manager.alloc_c128 raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") + def _scatter_c4_prefill_slots_slow(self, req_idx: int, first: int, last: int) -> None: + """Idempotence fallback for overlapped/repeated c4 prep.""" + page = DSV4_C4_PAGE_SIZE + mapping = self.mem_manager.full_to_c4_indexs + for page_base in range((first // page) * page, last, page): + e0 = max(first, page_base) + e1 = min(last, page_base + page) + entries = torch.arange(e0, e1, dtype=torch.long, device="cuda") + full_slots = self.req_to_token_indexs[req_idx, entries * 4 + 3].long() + existing = mapping[full_slots] + missing = existing < 0 + if not bool(missing.any()): + continue + + mapped = torch.nonzero(existing >= 0, as_tuple=False) + if mapped.numel() > 0: + j = int(mapped[0].item()) + base = int(existing[j].item()) - ((e0 + j) % page) + elif e0 > page_base: + prev_full = self.req_to_token_indexs[req_idx, e0 * 4 - 1].long() + prev_slot = int(mapping[prev_full].item()) + assert prev_slot >= 0 and prev_slot % page == (e0 - 1) % page + base = prev_slot - ((e0 - 1) % page) + else: + base = int(self.mem_manager.alloc_c4_pages(1)[0].item()) * page + + slots = (base + entries % page).to(torch.int32) + if mapped.numel() > 0: + assert bool((existing[existing >= 0] == slots[existing >= 0]).all()) + mapping[full_slots[missing]] = slots[missing] + self.mem_manager.count_c4_slots(slots[missing], 1) + return + + def _scatter_c4_prefill_slots(self, req_idx: int, first: int, last: int) -> None: + """为 logical c4 entry [first, last) 分配 page-safe c4 槽。 + + 不变式: logical entry e 映射到 physical_page * 64 + e % 64,同一 logical page + 内 entry 共享 physical_page。这是 DeepGEMM paged MQA logits 直接消费 page table 的前提。 + """ + if last <= first: + return + page = DSV4_C4_PAGE_SIZE + mapping = self.mem_manager.full_to_c4_indexs + entries = torch.arange(first, last, dtype=torch.long, device="cuda") + full_slots = self.req_to_token_indexs[req_idx, entries * 4 + 3].long() + need = mapping[full_slots] < 0 + if not bool(need.any()): + return + if not bool(need.all()): + self._scatter_c4_prefill_slots_slow(req_idx, first, last) + return + + first_page = first // page + last_page = (last - 1) // page + n_pages = last_page - first_page + 1 + bases = torch.empty((n_pages,), dtype=torch.long, device="cuda") + + base_start = 0 + if first % page != 0: + prev_full = self.req_to_token_indexs[req_idx, first * 4 - 1].long() + prev_slot = int(mapping[prev_full].item()) + assert prev_slot >= 0 and prev_slot % page == (first - 1) % page + bases[0] = prev_slot - ((first - 1) % page) + base_start = 1 + + new_page_count = n_pages - base_start + if new_page_count > 0: + new_pages = self.mem_manager.alloc_c4_pages(new_page_count).cuda(non_blocking=True).long() + bases[base_start:] = new_pages * page + + page_local = torch.div(entries, page, rounding_mode="floor") - first_page + slots = (bases[page_local] + entries % page).to(torch.int32) + mapping[full_slots] = slots + self.mem_manager.count_c4_slots(slots, 1) + return + + def _scatter_c4_decode_slots( + self, + b_req_idx_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, + mem_indexes: torch.Tensor, + ) -> None: + page = DSV4_C4_PAGE_SIZE + mapping = self.mem_manager.full_to_c4_indexs + req_list = b_req_idx_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() + mem_indexes = mem_indexes.cuda().long().reshape(-1) + + cont_rows, cont_prev_pos, cont_offsets = [], [], [] + new_rows = [] + for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): + req_idx, seq_len = int(req_idx), int(seq_len) + if req_idx == self.HOLD_REQUEST_ID or seq_len <= 0 or seq_len % 4 != 0: + continue + entry = seq_len // 4 - 1 + offset = entry % page + if offset == 0: + new_rows.append(i) + else: + cont_rows.append(i) + cont_prev_pos.append(entry * 4 - 1) + cont_offsets.append(offset) + + if cont_rows: + req_rows = torch.tensor([req_list[i] for i in cont_rows], dtype=torch.long, device="cuda") + prev_pos = torch.tensor(cont_prev_pos, dtype=torch.long, device="cuda") + prev_full = self.req_to_token_indexs[req_rows, prev_pos].long() + prev_slots = mapping[prev_full] + offsets = torch.tensor(cont_offsets, dtype=torch.int32, device="cuda") + assert bool((prev_slots >= 0).all()) + assert bool(((prev_slots % page) == (offsets - 1)).all()) + slots = (prev_slots + 1).to(torch.int32) + mapping[mem_indexes[cont_rows]] = slots + self.mem_manager.count_c4_slots(slots, 1) + + if new_rows: + pages = self.mem_manager.alloc_c4_pages(len(new_rows)).cuda(non_blocking=True).long() + slots = (pages * page).to(torch.int32) + mapping[mem_indexes[new_rows]] = slots + self.mem_manager.count_c4_slots(slots, 1) + return + def _scatter_compress_slots(self, ratio: int, full_slots: torch.Tensor) -> None: """为组末 full 槽位分配压缩槽并写入映射。已映射(>=0)的行跳过——重复 prep 幂等。""" if full_slots.numel() == 0: @@ -568,9 +694,15 @@ def prepare_prefill_compress_slots( req_list = b_req_idx.detach().cpu().tolist() ready_list = b_ready_cache_len.detach().cpu().tolist() seq_list = b_seq_len.detach().cpu().tolist() - for ratio, n_layers in ((4, self.n_c4), (128, self.n_c128)): - if n_layers == 0: - continue + if self.n_c4 > 0: + for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + self._scatter_c4_prefill_slots(req_idx, int(ready_len) // 4, int(seq_len) // 4) + + if self.n_c128 > 0: + ratio = 128 end_slots = [] for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): req_idx = int(req_idx) @@ -597,9 +729,11 @@ def prepare_decode_compress_slots( return req_list = b_req_idx_cpu.tolist() seq_list = b_seq_len_cpu.tolist() - for ratio, n_layers in ((4, self.n_c4), (128, self.n_c128)): - if n_layers == 0: - continue + if self.n_c4 > 0: + self._scatter_c4_decode_slots(b_req_idx_cpu, b_seq_len_cpu, mem_indexes) + + if self.n_c128 > 0: + ratio = 128 rows = [ i for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) @@ -624,7 +758,7 @@ def get_prompt_cache_value_ops(self): return DeepseekV4PromptCacheValueOps(self) def get_prompt_cache_page_size(self): - return 128 + return DSV4_PROMPT_CACHE_PAGE_SIZE def compute_swa_page_valid(self, full_slots: torch.Tensor) -> torch.Tensor: """按当下 full_to_swa 映射给出按页有效性: full_slots [L](L 为 page 整数倍) -> @@ -654,7 +788,7 @@ def slice_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload, star start = int(start) end = int(end) page = self.get_prompt_cache_page_size() - # radix page=128 保证分裂点页对齐,bitmap 可整页切分。 + # radix page 保证分裂点页对齐,bitmap 可整页切分。 return DeepseekV4PromptCachePayload( cache_len=end - start, swa_page_valid=payload.swa_page_valid[start // page : end // page].clone() diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 6ebc0b2856..2721c23ac4 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,3 +1,4 @@ +import os import torch import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl @@ -585,13 +586,26 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, ) return slots.unsqueeze(1), lengths - import deep_gemm - from ..triton_kernel.gather_c4_indexer_k_dsv4 import gather_c4_indexer_k_ragged - b_req_idx = infer_state.b_req_idx batch = b_req_idx.shape[0] device = positions.device c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req + + if os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER", "0") == "1": + out = self._c4_indices_paged( + infer_state=infer_state, + idx_q_fp8=idx_q_fp8, + weights=weights, + positions=positions, + c4_len=c4_len, + c4_cap=c4_cap, + ) + if out is not None: + return out + + import deep_gemm + from ..triton_kernel.gather_c4_indexer_k_dsv4 import gather_c4_indexer_k_ragged + k_fp8, k_scale, ragged_slots = gather_c4_indexer_k_ragged( mem_manager, self.layer_idx_, @@ -622,3 +636,72 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, top_slots = torch.where(invalid, torch.full_like(top_slots, -1), top_slots) topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) return top_slots.unsqueeze(1), topk_lengths.contiguous() + + def _c4_indices_paged(self, infer_state, idx_q_fp8, weights, positions, c4_len, c4_cap): + import deep_gemm + from sglang.jit_kernel.dsv4 import topk_transform_512 + from ..triton_kernel.gather_c4_indexer_k_dsv4 import build_c4_indexer_page_table + + mem_manager = infer_state.mem_manager + index_topk = self.index_topk + device = positions.device + b_req_idx = infer_state.b_req_idx + batch = b_req_idx.shape[0] + validate = os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER_VALIDATE", "0") == "1" + validate_now = validate and not torch.cuda.is_current_stream_capturing() + + page_table, valid_flag = build_c4_indexer_page_table( + mem_manager, + b_req_idx, + c4_len, + c4_cap, + infer_state.req_manager.req_to_token_indexs, + infer_state.req_manager.HOLD_REQUEST_ID, + validate=validate_now, + ) + if validate_now and int(valid_flag.item()) == 0: + if os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER_STRICT", "0") == "1": + raise RuntimeError("DeepSeek-V4 paged indexer requires page-aligned c4 slots") + return None + + if infer_state.is_prefill: + token_batch_pos = torch.repeat_interleave( + torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len + ) + row_page_table = page_table[token_batch_pos.long()].contiguous() + else: + row_page_table = page_table + + valid_len = ((positions + 1) // 4).to(torch.int32) + ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1).contiguous() + kv_cache = mem_manager.c4_indexer_pool.get_layer_buffer(mem_manager.layer_to_c4_idx[self.layer_idx_]).view( + mem_manager.c4_indexer_pool.num_pages, + mem_manager.c4_indexer_pool.page_size, + 1, + self.index_head_dim + 4, + ) + metadata = deep_gemm.get_paged_mqa_logits_metadata( + ctx_lens, + mem_manager.c4_indexer_pool.page_size, + deep_gemm.get_num_sms(), + ) + logits = deep_gemm.fp8_paged_mqa_logits( + idx_q_fp8.unsqueeze(1), + kv_cache, + weights, + ctx_lens, + row_page_table, + metadata, + c4_cap, + False, + ) + top_slots = torch.empty((idx_q_fp8.shape[0], index_topk), dtype=torch.int32, device=device) + topk_transform_512( + logits, + valid_len, + row_page_table, + top_slots, + mem_manager.c4_indexer_pool.page_size, + ) + topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) + return top_slots.unsqueeze(1), topk_lengths.contiguous() diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py index b6e6ba751a..a7a0a4be85 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -107,3 +107,94 @@ def gather_c4_indexer_k_ragged( num_warps=1, ) return k_fp8, k_scale, slots + + +@triton.jit +def _build_c4_indexer_page_table_kernel( + req_idx_ptr, # [batch] int + c4_len_ptr, # [batch] int + req_to_token_ptr, + req_to_token_stride0, + full_to_c4_ptr, + page_table_ptr, # [batch, page_cap] int32 + valid_flag_ptr, # [1] int32, initialized to 1; set to 0 on layout mismatch + page_cap, + hold_req_id, + RATIO: tl.constexpr, + PAGE_SIZE: tl.constexpr, + VALIDATE: tl.constexpr, +): + p = tl.program_id(0) + r = tl.program_id(1) + req = tl.load(req_idx_ptr + r).to(tl.int64) + c4_len = tl.load(c4_len_ptr + r).to(tl.int64) + page_start = p * PAGE_SIZE + active = (req != hold_req_id) & (page_start < c4_len) + + full_pos0 = page_start * RATIO + (RATIO - 1) + full_slot0 = tl.load( + req_to_token_ptr + req * req_to_token_stride0 + full_pos0, + mask=active, + other=0, + ).to(tl.int64) + c4_slot0 = tl.load(full_to_c4_ptr + full_slot0, mask=active, other=0).to(tl.int64) + phys_page = c4_slot0 // PAGE_SIZE + tl.store(page_table_ptr + r * page_cap + p, tl.where(active, phys_page, 0).to(tl.int32)) + + if VALIDATE: + offs = tl.arange(0, PAGE_SIZE) + e = page_start + offs + valid = active & (e < c4_len) + full_pos = e * RATIO + (RATIO - 1) + full_slot = tl.load( + req_to_token_ptr + req * req_to_token_stride0 + full_pos, + mask=valid, + other=0, + ).to(tl.int64) + c4_slot = tl.load(full_to_c4_ptr + full_slot, mask=valid, other=-1).to(tl.int64) + expected = phys_page * PAGE_SIZE + offs + ok = tl.where(valid, (c4_slot == expected) & (c4_slot >= 0), True) + if tl.min(ok.to(tl.int32), axis=0) == 0: + tl.store(valid_flag_ptr, 0) + + +@torch.no_grad() +def build_c4_indexer_page_table( + mem_manager, + b_req_idx: torch.Tensor, + c4_len: torch.Tensor, + c4_cap: int, + req_to_token_indexs: torch.Tensor, + hold_req_id: int, + validate: bool = False, +): + """Build the logical-c4-page -> physical-c4-page table expected by DeepGEMM paged logits. + + This is safe only when each logical c4 page maps to a physical page with matching offsets: + c4_slot(entry p*64 + o) == page_table[p] * 64 + o + The optional validation flag checks that invariant and lets the caller fall back to the + gather path while we keep the current token-slot allocator. + """ + pool = mem_manager.c4_indexer_pool + page_size = pool.page_size + assert c4_cap % page_size == 0 + batch = b_req_idx.shape[0] + page_cap = c4_cap // page_size + page_table = torch.empty((batch, page_cap), dtype=torch.int32, device=b_req_idx.device) + valid_flag = torch.ones((1,), dtype=torch.int32, device=b_req_idx.device) + _build_c4_indexer_page_table_kernel[(page_cap, batch)]( + b_req_idx, + c4_len, + req_to_token_indexs, + req_to_token_indexs.stride(0), + mem_manager.full_to_c4_indexs, + page_table, + valid_flag, + page_cap, + int(hold_req_id), + RATIO=4, + PAGE_SIZE=page_size, + VALIDATE=validate, + num_warps=1, + ) + return page_table, valid_flag diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index a852461df0..b667c4be72 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -8,7 +8,7 @@ from sortedcontainers import SortedDict from dataclasses import dataclass, field from typing import List, Dict, Tuple, Optional, Callable, Any, Union -from lightllm.common.req_manager import ReqManager, ReqManagerForMamba +from lightllm.common.req_manager import DeepseekV4ReqManager, ReqManager, ReqManagerForMamba from lightllm.utils.infer_utils import mark_start, mark_end from lightllm.server.core.objs import Req, SamplingParams, FinishStatus, ShmReqManager from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache, TreeNode @@ -177,6 +177,7 @@ def _dsv4_full_att_free_req(self, free_token_index: List, req: "InferReq"): # 载荷只剩按页 bitmap(compressor 状态随 swa 页生灭/边界自然归零,不进载荷), # 任意 128 对齐前缀皆可插入——含生成段(floor(cur_kv_len) 边界,回收保留尾页保证其驻留)。 cache_len = self.radix_cache.align_len(req.cur_kv_len) + self.req_manager: DeepseekV4ReqManager if cache_len > old_prefix_len: payload = self.req_manager.build_prompt_cache_payload(req.req_idx, cache_len) value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() From 52a1528051a53f0f1291332357d7df5730a5df59 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 18 Jun 2026 03:07:49 +0000 Subject: [PATCH 028/214] fix chunk_size and page_size --- lightllm/server/router/model_infer/infer_batch.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index b667c4be72..2bf2314185 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -916,10 +916,16 @@ def _align_chuncked_end_for_prompt_cache(self, chunked_start: int, chunked_end: page_size = getattr(radix_cache, "page_size", 1) if radix_cache is not None else 1 if page_size <= 1 or self.sampling_param.disable_prompt_cache: return chunked_end - prompt_end = self.shm_req.input_len - next_page_end = ((int(chunked_start) // page_size) + 1) * page_size - if int(chunked_start) < next_page_end < int(chunked_end) and next_page_end <= prompt_end: - return next_page_end + prompt_end = int(self.shm_req.input_len) + chunked_start = int(chunked_start) + chunked_end = int(chunked_end) + if chunked_end >= prompt_end: + return chunked_end + + assert self.args.chunked_prefill_size % page_size == 0, ( + f"chunked_prefill_size={self.args.chunked_prefill_size} must be divisible by " + f"prompt-cache page_size={page_size}" + ) return chunked_end def get_chuncked_input_token_len_for_linear_att(self): From 0dbc90b6a233ffe4dfbea2c4a237b5a18f88a612 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 18 Jun 2026 03:08:41 +0000 Subject: [PATCH 029/214] add sglang third_party --- .../layer_infer/transformer_layer_infer.py | 4 +- lightllm/third_party/__init__.py | 1 + lightllm/third_party/sglang_jit/LICENSE | 201 ++++ lightllm/third_party/sglang_jit/README.md | 13 + lightllm/third_party/sglang_jit/__init__.py | 1 + .../sglang_jit/csrc/deepseek_v4/c128.cuh | 522 +++++++++++ .../csrc/deepseek_v4/c128_online.cuh | 726 +++++++++++++++ .../csrc/deepseek_v4/c128_online_v2.cuh | 875 ++++++++++++++++++ .../sglang_jit/csrc/deepseek_v4/c128_v2.cuh | 448 +++++++++ .../sglang_jit/csrc/deepseek_v4/c4.cuh | 549 +++++++++++ .../sglang_jit/csrc/deepseek_v4/c4_v2.cuh | 405 ++++++++ .../sglang_jit/csrc/deepseek_v4/c_plan.cuh | 839 +++++++++++++++++ .../sglang_jit/csrc/deepseek_v4/common.cuh | 208 +++++ .../csrc/deepseek_v4/fused_norm_rope.cuh | 254 +++++ .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 643 +++++++++++++ .../sglang_jit/csrc/deepseek_v4/hash_topk.cuh | 214 +++++ .../csrc/deepseek_v4/hisparse_transfer.cuh | 82 ++ .../csrc/deepseek_v4/main_norm_rope.cuh | 845 +++++++++++++++++ .../deepseek_v4/mega_moe_pre_dispatch.cuh | 219 +++++ .../csrc/deepseek_v4/paged_mqa_metadata.cuh | 119 +++ .../sglang_jit/csrc/deepseek_v4/rope.cuh | 169 ++++ .../silu_and_mul_masked_post_quant.cuh | 540 +++++++++++ .../sglang_jit/csrc/deepseek_v4/store.cuh | 205 ++++ .../sglang_jit/csrc/deepseek_v4/topk_v1.cuh | 340 +++++++ .../sglang_jit/csrc/deepseek_v4/topk_v2.cuh | 493 ++++++++++ .../third_party/sglang_jit/dsv4/__init__.py | 8 + .../sglang_jit/dsv4/elementwise.py | 215 +++++ lightllm/third_party/sglang_jit/dsv4/topk.py | 92 ++ lightllm/third_party/sglang_jit/dsv4/utils.py | 2 + .../sglang_jit/include/sgl_kernel/atomic.cuh | 35 + .../sglang_jit/include/sgl_kernel/cta.cuh | 40 + .../sgl_kernel/deepseek_v4/compress.cuh | 37 + .../sgl_kernel/deepseek_v4/compress_v2.cuh | 99 ++ .../sgl_kernel/deepseek_v4/fp8_utils.cuh | 112 +++ .../sgl_kernel/deepseek_v4/kvcacheio.cuh | 96 ++ .../sgl_kernel/deepseek_v4/topk/cluster.cuh | 257 +++++ .../sgl_kernel/deepseek_v4/topk/common.cuh | 176 ++++ .../sgl_kernel/deepseek_v4/topk/ptx.cuh | 54 ++ .../sgl_kernel/deepseek_v4/topk/register.cuh | 302 ++++++ .../sgl_kernel/deepseek_v4/topk/streaming.cuh | 213 +++++ .../include/sgl_kernel/distributed/common.cuh | 120 +++ .../distributed/custom_all_reduce.cuh | 354 +++++++ .../sglang_jit/include/sgl_kernel/ffi.h | 104 +++ .../include/sgl_kernel/impl/norm.cuh | 168 ++++ .../sglang_jit/include/sgl_kernel/math.cuh | 71 ++ .../sglang_jit/include/sgl_kernel/runtime.cuh | 86 ++ .../include/sgl_kernel/scalar_type.hpp | 334 +++++++ .../include/sgl_kernel/source_location.h | 40 + .../sglang_jit/include/sgl_kernel/tensor.h | 605 ++++++++++++ .../sglang_jit/include/sgl_kernel/tile.cuh | 62 ++ .../sglang_jit/include/sgl_kernel/type.cuh | 120 +++ .../sglang_jit/include/sgl_kernel/utils.cuh | 333 +++++++ .../sglang_jit/include/sgl_kernel/utils.h | 186 ++++ .../sglang_jit/include/sgl_kernel/vec.cuh | 118 +++ .../sglang_jit/include/sgl_kernel/warp.cuh | 56 ++ lightllm/third_party/sglang_jit/jit_utils.py | 432 +++++++++ .../third_party/sglang_jit/runtime_utils.py | 5 + 57 files changed, 13845 insertions(+), 2 deletions(-) create mode 100644 lightllm/third_party/__init__.py create mode 100755 lightllm/third_party/sglang_jit/LICENSE create mode 100644 lightllm/third_party/sglang_jit/README.md create mode 100644 lightllm/third_party/sglang_jit/__init__.py create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh create mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/dsv4/__init__.py create mode 100644 lightllm/third_party/sglang_jit/dsv4/elementwise.py create mode 100644 lightllm/third_party/sglang_jit/dsv4/topk.py create mode 100644 lightllm/third_party/sglang_jit/dsv4/utils.py create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh create mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh create mode 100644 lightllm/third_party/sglang_jit/jit_utils.py create mode 100644 lightllm/third_party/sglang_jit/runtime_utils.py diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 2721c23ac4..617d0dcd85 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -144,7 +144,7 @@ def _get_qkv( infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, ): - from sglang.jit_kernel.dsv4 import fused_q_norm_rope + from lightllm.third_party.sglang_jit.dsv4 import fused_q_norm_rope input = self._tpsp_allgather(input=input, infer_state=infer_state) T = input.shape[0] @@ -639,7 +639,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, def _c4_indices_paged(self, infer_state, idx_q_fp8, weights, positions, c4_len, c4_cap): import deep_gemm - from sglang.jit_kernel.dsv4 import topk_transform_512 + from lightllm.third_party.sglang_jit.dsv4 import topk_transform_512 from ..triton_kernel.gather_c4_indexer_k_dsv4 import build_c4_indexer_page_table mem_manager = infer_state.mem_manager diff --git a/lightllm/third_party/__init__.py b/lightllm/third_party/__init__.py new file mode 100644 index 0000000000..2adb50db25 --- /dev/null +++ b/lightllm/third_party/__init__.py @@ -0,0 +1 @@ +"""Third-party source subsets vendored for LightLLM runtime support.""" diff --git a/lightllm/third_party/sglang_jit/LICENSE b/lightllm/third_party/sglang_jit/LICENSE new file mode 100755 index 0000000000..9c422689c8 --- /dev/null +++ b/lightllm/third_party/sglang_jit/LICENSE @@ -0,0 +1,201 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2023-2024 SGLang Team + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/lightllm/third_party/sglang_jit/README.md b/lightllm/third_party/sglang_jit/README.md new file mode 100644 index 0000000000..4f68c9cfd8 --- /dev/null +++ b/lightllm/third_party/sglang_jit/README.md @@ -0,0 +1,13 @@ +# Vendored SGLang JIT Subset + +This directory contains the minimal SGLang JIT source subset needed by the +DeepSeek-V4 LightLLM implementation. + +Source: https://github.com/sgl-project/sglang +Commit: 8cea0473ea5299bc04885f8f6ba71269415a39b5 +License: Apache License 2.0, copied in `LICENSE`. + +Local changes: +- The Python imports were moved from `sglang.jit_kernel.*` to + `lightllm.third_party.sglang_jit.*`. +- The package exports only the DSv4 functions used by LightLLM. diff --git a/lightllm/third_party/sglang_jit/__init__.py b/lightllm/third_party/sglang_jit/__init__.py new file mode 100644 index 0000000000..164d545b4e --- /dev/null +++ b/lightllm/third_party/sglang_jit/__init__.py @@ -0,0 +1 @@ +"""Vendored SGLang JIT kernels used by DeepSeek-V4.""" diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh new file mode 100644 index 0000000000..3a89e8114c --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh @@ -0,0 +1,522 @@ +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include + +namespace { + +using Plan128 = device::compress::PrefillPlan; +using IndiceT = int32_t; + +/// \brief Each thread will handle this many elements (split along head_dim) +constexpr int32_t kTileElements = 2; +/// \brief Each warp will handle this many elements (split along 128) +constexpr int32_t kElementsPerWarp = 8; +constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; +constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps; + +/// \brief Need to reduce register usage to increase occupancy +#define C128_KERNEL __global__ __launch_bounds__(kBlockSize, 2) + +struct Compress128DecodeParams { + /** + * \brief Shape: `[num_indices, 128, head_dim * 2]` \n + * last dimension layout: + * | kv current | score current | + */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[batch_size, head_dim * 2]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[batch_size, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[128, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, ]` */ + const IndiceT* __restrict__ seq_lens; + /** \NOTE: `batch_size` <= `num_indices` */ + uint32_t batch_size; +}; + +struct Compress128PrefillParams { + /** + * \brief Shape: `[num_indices, 128, head_dim * 2]` \n + * last dimension layout: + * | kv current | score current | + */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[batch_size, head_dim * 2]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[batch_size, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[128, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, ]`*/ + const int32_t* __restrict__ load_indices; + /** \brief The following part is plan info. */ + const Plan128* __restrict__ compress_plan; + const Plan128* __restrict__ write_plan; + uint32_t num_compress; + uint32_t num_write; +}; + +struct Compress128SharedBuffer { + using Storage = device::AlignedVector; + Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict + SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { + return data[warp_id][lane_id]; + } + SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { + return data[warp_id][lane_id][tile_id]; + } +}; + +template +SGL_DEVICE void c128_write( + T* kv_score_buf, // + const T* kv_score_src, + const int64_t head_dim, + const int32_t write_pos, + const uint32_t lane_id) { + using namespace device; + + using Storage = AlignedVector; + const auto element_size = head_dim * 2; + const auto gmem = tile::Memory{lane_id, kWarpThreads}; + kv_score_buf += write_pos * element_size; + + /// NOTE: Layout | [0] = kv | [1] = score | + Storage kv_score[2]; +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + kv_score[i] = gmem.load(kv_score_src + head_dim * i); + } +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + gmem.store(kv_score_buf + head_dim * i, kv_score[i]); + } +} + +template +SGL_DEVICE void c128_forward( + const InFloat* kv_score_buf, + const InFloat* kv_score_src, + OutFloat* kv_out, + const InFloat* score_bias, + const int64_t head_dim, + const int32_t window_len, + const uint32_t warp_id, + const uint32_t lane_id) { + using namespace device; + + const auto element_size = head_dim * 2; + const auto score_offset = head_dim; + + /// NOTE: part 1: load kv + score + using StorageIn = AlignedVector; + const auto gmem_in = tile::Memory{lane_id, kWarpThreads}; + StorageIn kv[kElementsPerWarp]; + StorageIn score[kElementsPerWarp]; + StorageIn bias[kElementsPerWarp]; + const int32_t warp_offset = warp_id * kElementsPerWarp; + +#pragma unroll + for (int32_t i = 0; i < 8; ++i) { + const int32_t j = i + warp_offset; + bias[i] = gmem_in.load(score_bias + j * head_dim); + } + +#pragma unroll + for (int32_t i = 0; i < kElementsPerWarp; ++i) { + const int32_t j = i + warp_offset; + const InFloat* src; + __builtin_assume(j < 128); + if (j < window_len) { + src = kv_score_buf + j * element_size; + } else { + /// NOTE: k in [-127, 0]. We'll load from the ragged `kv_score_src` + const int32_t k = j - 127; + src = kv_score_src + k * element_size; + } + kv[i] = gmem_in.load(src); + score[i] = gmem_in.load(src + score_offset); + } + + /// NOTE: part 2: safe online softmax + weighted sum + using TmpStorage = typename Compress128SharedBuffer::Storage; + __shared__ Compress128SharedBuffer s_local_val_max; + __shared__ Compress128SharedBuffer s_local_exp_sum; + __shared__ Compress128SharedBuffer s_local_product; + + TmpStorage tmp_val_max; + TmpStorage tmp_exp_sum; + TmpStorage tmp_product; + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + float score_fp32[kElementsPerWarp]; + +#pragma unroll + for (int32_t j = 0; j < kElementsPerWarp; ++j) { + score_fp32[j] = cast(score[j][i]) + cast(bias[j][i]); + } + + float max_value = score_fp32[0]; + float sum_exp_value = 0.0f; + +#pragma unroll + for (int32_t j = 1; j < kElementsPerWarp; ++j) { + const auto fp32_score = score_fp32[j]; + max_value = fmaxf(max_value, fp32_score); + } + + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + const auto fp32_score = score_fp32[j]; + const auto exp_score = expf(fp32_score - max_value); + sum_product += cast(kv[j][i]) * exp_score; + sum_exp_value += exp_score; + } + + tmp_val_max[i] = max_value; + tmp_exp_sum[i] = sum_exp_value; + tmp_product[i] = sum_product; + } + + // naturally aligned, so no bank conflict + s_local_val_max(warp_id, lane_id) = tmp_val_max; + s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; + s_local_product(warp_id, lane_id) = tmp_product; + + __syncthreads(); + + /// NOTE: part 3: online softmax + /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce + /// each reduce will consume `kNumWarps` threads (use partial warp reduction) + constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; + constexpr uint32_t kIteration = kReductionCount / kBlockSize; + +#pragma unroll + for (uint32_t i = 0; i < kIteration; ++i) { + /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` + const uint32_t j = i * kBlockSize + warp_id * kWarpThreads + lane_id; + /// NOTE: Range `[0, kNumWarps)` + const uint32_t local_warp_id = j % kNumWarps; + /// NOTE: Range `[0, kTileElements * kWarpThreads)` + const uint32_t local_elem_id = j / kNumWarps; + /// NOTE: Range `[0, kTileElements)` + const uint32_t local_tile_id = local_elem_id % kTileElements; + /// NOTE: Range `[0, kWarpThreads)` + const uint32_t local_lane_id = local_elem_id / kTileElements; + /// NOTE: each warp will access the whole tile (all `kTileElements`) + /// and for different lanes, the memory access only differ in `local_warp_id` + /// so there's no bank conflict in shared memory access. + static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); + const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); + const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); + const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); + const auto global_val_max = warp::reduce_max(local_val_max); + const auto rescale = expf(local_val_max - global_val_max); + const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); + const auto final_scale = rescale / global_exp_sum; + const auto global_product = warp::reduce_sum(local_product * final_scale); + kv_out[local_elem_id] = cast(global_product); + } +} + +template +C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + constexpr int64_t kElementSize = kHeadDim * 2; + static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); + + const auto& [ + _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score + indices, seq_lens, batch_size // decode info + ] = params; + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + + const uint32_t global_bid = blockIdx.x / kNumSplit; // batch id + const uint32_t global_sid = blockIdx.x % kNumSplit; // split id + if (global_bid >= batch_size) return; + + const int32_t index = indices[global_bid]; + const int32_t seq_len = seq_lens[global_bid]; + const int64_t split_offset = global_sid * kTileDim; + + // kv score + const auto kv_score_buffer = static_cast(_kv_score_buffer); + const auto kv_buf = kv_score_buffer + index * (kElementSize * 128) + split_offset; + + // kv input + const auto kv_score_input = static_cast(_kv_score_input); + const auto kv_src = kv_score_input + global_bid * kElementSize + split_offset; + + // kv output + const auto kv_compressed_output = static_cast(_kv_compressed_output); + const auto kv_out = kv_compressed_output + global_bid * kHeadDim + split_offset; + + // score bias (ape) + const auto score_bias = static_cast(_score_bias) + split_offset; + + PDLWaitPrimary(); + + /// NOTE: the write must be visible to the subsequent c128_forward, + /// so only the last warp can write to HBM + /// In addition, `position` = `seq_len - 1`. To avoid underflow, we use `seq_len + 127` + if (warp_id == kNumWarps - 1) { + c128_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 127) % 128, lane_id); + } + if (seq_len % 128 == 0) { + c128_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, /*window_len=*/128, warp_id, lane_id); + } + + PDLTriggerSecondary(); +} + +// compress kernel +template +C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + constexpr int64_t kElementSize = kHeadDim * 2; + static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); + + const auto& [ + _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score + indices, load_indices, compress_plan, write_plan, num_compress, num_write // prefill plan + ] = params; + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + + uint32_t global_id; + if constexpr (kWrite) { + // for write kernel, we use global warp_id to dispatch work + global_id = (blockIdx.x * blockDim.x + threadIdx.x) / kWarpThreads; + } else { + // for compress kernel, we use block id to dispatch work + global_id = blockIdx.x; // block id + } + const uint32_t global_pid = global_id / kNumSplit; // plan id + const uint32_t global_sid = global_id % kNumSplit; // split id + + /// NOTE: compiler can optimize this if-else at compile time + const auto num_plans = kWrite ? num_write : num_compress; + const auto plan_ptr = kWrite ? write_plan : compress_plan; + if (global_pid >= num_plans) return; + + const auto& [ragged_id, global_bid, position, window_len] = plan_ptr[global_pid]; + const auto indices_ptr = kWrite ? indices : load_indices; + + const int64_t split_offset = global_sid * kTileDim; + + // kv input + const auto kv_score_input = static_cast(_kv_score_input); + const auto kv_src = kv_score_input + ragged_id * kElementSize + split_offset; + + // kv output + const auto kv_compressed_output = static_cast(_kv_compressed_output); + const auto kv_out = kv_compressed_output + ragged_id * kHeadDim + split_offset; + + // score bias (ape) + const auto score_bias = static_cast(_score_bias) + split_offset; + + if (ragged_id == 0xFFFFFFFF) [[unlikely]] + return; + + const int32_t index = indices_ptr[global_bid]; + // kv score + const auto kv_score_buffer = static_cast(_kv_score_buffer); + const auto kv_buf = kv_score_buffer + index * (kElementSize * 128) + split_offset; + + PDLWaitPrimary(); + + // only responsible for the compress part + if constexpr (kWrite) { + c128_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 128, lane_id); + } else { + c128_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, window_len, warp_id, lane_id); + } + + PDLTriggerSecondary(); +} + +template +struct FlashCompress128Kernel { + static constexpr auto decode_kernel = flash_c128_decode; + template + static constexpr auto prefill_kernel = flash_c128_prefill; + static constexpr auto prefill_c_kernel = prefill_kernel; + static constexpr auto prefill_w_kernel = prefill_kernel; + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static constexpr uint32_t kWriteBlockSize = 128; + static constexpr uint32_t kWarpsPerWriteBlock = kWriteBlockSize / device::kWarpThreads; + + static void run_decode( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::Optional /* UNUSED */) { + using namespace host; + + // this should not happen in practice + auto B = SymbolicSize{"batch_size"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({-1, 128, kHeadDim * 2}) // kv score + .with_dtype() + .with_device(device) + .verify(kv_score_buffer); + TensorMatcher({B, kHeadDim * 2}) // kv score input + .with_dtype() + .with_device(device) + .verify(kv_score_input); + TensorMatcher({B, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device) + .verify(ape); + TensorMatcher({B}) // indices + .with_dtype() + .with_device(device) + .verify(indices); + TensorMatcher({B}) // seq lens + .with_dtype() + .with_device(device) + .verify(seq_lens); + + const auto batch_size = static_cast(B.unwrap()); + const auto params = Compress128DecodeParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .batch_size = batch_size, + }; + + const uint32_t num_blocks = batch_size * kNumSplit; + LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // + .enable_pdl(kUsePDL)(decode_kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView compress_plan, + const tvm::ffi::TensorView write_plan, + const tvm::ffi::Optional extra) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto N = SymbolicSize{"num_q_tokens"}; + auto X = SymbolicSize{"compress_tokens"}; + auto Y = SymbolicSize{"write_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 128, kHeadDim * 2}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({N, kHeadDim * 2}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({N, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + TensorMatcher({B}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + TensorMatcher({X, compress::kPrefillPlanDim}) // compress plan + .with_dtype() + .with_device(device_) + .verify(compress_plan); + TensorMatcher({Y, compress::kPrefillPlanDim}) // write plan + .with_dtype() + .with_device(device_) + .verify(write_plan); + + // might be needed for prefill write + const auto load_indices = extra.value_or(indices); + TensorMatcher({B}) // [read_positions] + .with_dtype() + .with_device(device_) + .verify(load_indices); + + const auto device = device_.unwrap(); + const auto batch_size = static_cast(B.unwrap()); + const auto num_q_tokens = static_cast(N.unwrap()); + const auto num_c = static_cast(X.unwrap()); + const auto num_w = static_cast(Y.unwrap()); + const auto params = Compress128PrefillParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .load_indices = static_cast(load_indices.data_ptr()), + .compress_plan = static_cast(compress_plan.data_ptr()), + .write_plan = static_cast(write_plan.data_ptr()), + .num_compress = num_c, + .num_write = num_w, + }; + RuntimeCheck(num_q_tokens >= batch_size, "num_q_tokens must be >= batch_size"); + RuntimeCheck(num_q_tokens >= std::max(num_c, num_w), "invalid prefill plan"); + + constexpr auto kBlockSize_C = kBlockSize; + constexpr auto kBlockSize_W = kWriteBlockSize; + if (const auto num_c_blocks = num_c * kNumSplit) { + LaunchKernel(num_c_blocks, kBlockSize_C, device) // + .enable_pdl(kUsePDL)(prefill_c_kernel, params); + } + if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerWriteBlock)) { + LaunchKernel(num_w_blocks, kBlockSize_W, device) // + .enable_pdl(kUsePDL)(prefill_w_kernel, params); + } + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh new file mode 100644 index 0000000000..b497470606 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh @@ -0,0 +1,726 @@ +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include +#include +#include +#include + +#include +#include +#include + +namespace device::compress { + +/// \brief Plan entry for online compress 128 prefill. +/// Each entry describes a contiguous segment of tokens that lies inside a +/// single 128-chunk. Multiple segments can map to the same batch id when the +/// extend tokens span chunk boundaries. +/// +/// **Layout compatibility:** the field order/types match `PrefillPlan` so that +/// downstream kernels (e.g. `fused_norm_rope` in `CompressExtend` mode) can +/// consume the compress_plan tensor as-if it were a `PrefillPlan` tensor -- +/// they only read `ragged_id` and `position`, both of which carry identical +/// semantics here (the LAST token of the segment in q-ragged and global +/// coordinates respectively). +/// +/// Note that `window_len` here means "number of real tokens in this segment" +/// (1..128), which differs from `PrefillPlan::window_len`. Downstream kernels +/// that share the tensor MUST NOT read it under that name. +struct alignas(16) OnlinePrefillPlan { + /// \brief Ragged-q position of the LAST token in this segment. + /// Equal to `segment_start_ragged + window_len - 1`. + uint32_t ragged_id; + /// \brief Index into the `indices` / `load_indices` arrays. + uint32_t batch_id; + /// \brief Global position of the LAST token in this segment. + /// For compress plans, `position % 128 == 127` (chunk-closing); for write + /// plans, `position % 128 < 127`. + uint32_t position; + /// \brief Number of real tokens in this segment (1..128). + /// The first segment token sits at `position - window_len + 1` (global) and + /// at `ragged_id - window_len + 1` (ragged). + uint32_t window_len; +}; + +static_assert(alignof(OnlinePrefillPlan) == alignof(PrefillPlan)); +static_assert(sizeof(OnlinePrefillPlan) == sizeof(PrefillPlan)); + +} // namespace device::compress + +namespace host::compress { + +using device::compress::OnlinePrefillPlan; +using OnlinePrefillPlanTensorDtype = uint8_t; +inline constexpr int64_t kOnlinePrefillPlanDim = 16; + +static_assert(alignof(OnlinePrefillPlan) == sizeof(OnlinePrefillPlan)); +static_assert(sizeof(OnlinePrefillPlan) == kOnlinePrefillPlanDim * sizeof(OnlinePrefillPlanTensorDtype)); + +} // namespace host::compress + +namespace { + +using OnlinePlan = device::compress::OnlinePrefillPlan; +using IndiceT = int32_t; + +/// \brief Need to reduce register usage to increase occupancy +struct Compress128OnlineDecodeParams { + /** \brief Shape: `[num_indices, 1, head_dim * 3 (max, sum, kv) ]` \n */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[batch_size, head_dim * 2]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[batch_size, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[128, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, ]` */ + const IndiceT* __restrict__ seq_lens; + /** \NOTE: `batch_size` <= `num_indices` */ + uint32_t batch_size; +}; + +/// \brief Need to reduce register usage to increase occupancy +struct Compress128OnlinePrefillParams { + /** \brief Shape: `[num_indices, 1, head_dim * 3 (max, sum, kv) ]` \n */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[num_q_tokens, head_dim * 2]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[num_q_tokens, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[128, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ load_indices; + /// \brief Plan for segments that close a chunk (write to `kv_compressed_output`). + /// Shape: `[num_compress, 16]` (uint8). + const OnlinePlan* __restrict__ compress_plan; + /// \brief Plan for the trailing partial segment of each batch (write back to + /// `kv_score_buffer`). Shape: `[num_write, 16]` (uint8). + const OnlinePlan* __restrict__ write_plan; + uint32_t num_compress; + uint32_t num_write; +}; + +// 4 elements per thread, kHeadDim / 4 threads per block +template +__global__ void flash_c128_online_decode(const __grid_constant__ Compress128OnlineDecodeParams params) { + using namespace device; + constexpr uint32_t kVecSize = 4; + constexpr uint32_t kBlockSize = kHeadDim / kVecSize; + using Vec = AlignedVector; + const auto gmem = tile::Memory::cta(kBlockSize); + const auto batch_id = blockIdx.x; + const auto index = params.indices[batch_id]; + const auto seq_len = params.seq_lens[batch_id]; + + const auto kv_score_buffer = static_cast(params.kv_score_buffer); + const auto kv_buf = kv_score_buffer + index * (kHeadDim * 3); + const auto kv_score_input = static_cast(params.kv_score_input); + const auto kv_src = kv_score_input + batch_id * (kHeadDim * 2); + + /// NOTE: kv_score_buffer layout is [max, sum, kv] (slot 0 / 1 / 2). Reads, + /// writes, and the prefill kernel must all agree on this order. + const auto max_score_vec = gmem.load(kv_buf, 0); + const auto sum_score_vec = gmem.load(kv_buf, 1); + const auto old_kv_vec = gmem.load(kv_buf, 2); + + /// NOTE: kv_score_input layout is | kv | score | (head_dim each), matching + /// the offline c128 kernel and the online prefill kernel. + const auto new_kv_vec = gmem.load(kv_src, 0); + const auto new_score_raw_vec = gmem.load(kv_src, 1); + + /// NOTE: the new token sits at global position `seq_len - 1`, so its + /// position inside the 128-chunk is `(seq_len - 1) % 128`. The previous + /// `seq_len % 128` was off by one (`bias[127]` vs `bias[0]`, etc.). + const auto pos_in_chunk = (seq_len - 1) % 128; + const auto bias_vec = gmem.load(params.score_bias, pos_in_chunk); + + Vec out_kv_vec; + Vec out_max_vec; + Vec out_sum_vec; + if (pos_in_chunk != 0) { + // Mid-chunk: combine prior partial state with the new token via online softmax. +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const auto old_max = max_score_vec[i]; + const auto old_kv = old_kv_vec[i]; + const auto new_score = new_score_raw_vec[i] + bias_vec[i]; + const auto new_kv = new_kv_vec[i]; + const auto new_max = fmax(old_max, new_score); + const auto old_sum = sum_score_vec[i] * expf(old_max - new_max); + const auto new_exp = expf(new_score - new_max); + const auto new_sum = old_sum + new_exp; + out_kv_vec[i] = (old_kv * old_sum + new_kv * new_exp) / new_sum; + out_max_vec[i] = new_max; + out_sum_vec[i] = new_sum; + } + } else { + // First token of a new 128-chunk: initialize state with this token alone. +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + out_kv_vec[i] = new_kv_vec[i]; + out_max_vec[i] = new_score_raw_vec[i] + bias_vec[i]; + out_sum_vec[i] = 1.0f; // exp(score - max) with max == score + } + } + + if (pos_in_chunk == 127) { + // Chunk just closed: emit the compressed kv. No need to update the buffer + // -- the next chunk's first token will overwrite it. + const auto kv_out = static_cast(params.kv_compressed_output) + batch_id * kHeadDim; + gmem.store(kv_out, out_kv_vec); + } else { + // Otherwise persist the running [max, sum, kv] state for the next step. + gmem.store(kv_buf, out_max_vec, 0); + gmem.store(kv_buf, out_sum_vec, 1); + gmem.store(kv_buf, out_kv_vec, 2); + } +} + +constexpr int32_t kTileElements = 2; // split (along head-dim) +/// \brief Each warp will handle this many elements (split along softmax-128) +constexpr int32_t kElementsPerWarp = 8; +constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; +constexpr uint32_t kPrefillBlockSize = device::kWarpThreads * kNumWarps; +using PrefillStorage = device::AlignedVector; + +struct Compress128SharedBuffer { + using Storage = device::AlignedVector; + Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict + SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { + return data[warp_id][lane_id]; + } + SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { + return data[warp_id][lane_id][tile_id]; + } +}; + +template +SGL_DEVICE void c128_prefill_forward( + const PrefillStorage (&kv)[kElementsPerWarp], + const PrefillStorage (&score)[kElementsPerWarp], + float* kv_out, + float* max_out, + float* sum_out, + const uint32_t warp_id, + const uint32_t lane_id) { + using namespace device; + + /// NOTE: part 2: safe online softmax + weighted sum + using TmpStorage = typename Compress128SharedBuffer::Storage; + __shared__ Compress128SharedBuffer s_local_val_max; + __shared__ Compress128SharedBuffer s_local_exp_sum; + __shared__ Compress128SharedBuffer s_local_product; + + TmpStorage tmp_val_max; + TmpStorage tmp_exp_sum; + TmpStorage tmp_product; + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + float score_fp32[kElementsPerWarp]; + +#pragma unroll + for (int32_t j = 0; j < kElementsPerWarp; ++j) { + score_fp32[j] = score[j][i]; + } + + float max_value = score_fp32[0]; + float sum_exp_value = 0.0f; + +#pragma unroll + for (int32_t j = 1; j < kElementsPerWarp; ++j) { + const auto fp32_score = score_fp32[j]; + max_value = fmaxf(max_value, fp32_score); + } + + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + const auto fp32_score = score_fp32[j]; + const auto exp_score = expf(fp32_score - max_value); + sum_product += cast(kv[j][i]) * exp_score; + sum_exp_value += exp_score; + } + + tmp_val_max[i] = max_value; + tmp_exp_sum[i] = sum_exp_value; + tmp_product[i] = sum_product; + } + + // naturally aligned, so no bank conflict + s_local_val_max(warp_id, lane_id) = tmp_val_max; + s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; + s_local_product(warp_id, lane_id) = tmp_product; + + __syncthreads(); + + /// NOTE: part 3: online softmax + /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce + /// each reduce will consume `kNumWarps` threads (use partial warp reduction) + constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; + constexpr uint32_t kIteration = kReductionCount / kPrefillBlockSize; + +#pragma unroll + for (uint32_t i = 0; i < kIteration; ++i) { + /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` + const uint32_t j = i * kPrefillBlockSize + warp_id * kWarpThreads + lane_id; + /// NOTE: Range `[0, kNumWarps)` + const uint32_t local_warp_id = j % kNumWarps; + /// NOTE: Range `[0, kTileElements * kWarpThreads)` + const uint32_t local_elem_id = j / kNumWarps; + /// NOTE: Range `[0, kTileElements)` + const uint32_t local_tile_id = local_elem_id % kTileElements; + /// NOTE: Range `[0, kWarpThreads)` + const uint32_t local_lane_id = local_elem_id / kTileElements; + /// NOTE: each warp will access the whole tile (all `kTileElements`) + /// and for different lanes, the memory access only differ in `local_warp_id` + /// so there's no bank conflict in shared memory access. + static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); + const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); + const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); + const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); + const auto global_val_max = warp::reduce_max(local_val_max); + const auto rescale = expf(local_val_max - global_val_max); + const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); + const auto final_scale = rescale / global_exp_sum; + const auto global_product = warp::reduce_sum(local_product * final_scale); + kv_out[local_elem_id] = global_product; + if constexpr (kNeedData) { + max_out[local_elem_id] = global_val_max; + sum_out[local_elem_id] = global_exp_sum; + } + } + if constexpr (kNeedData) __syncthreads(); +} + +/// \brief Sentinel score for padded positions in a 128-segment. +/// Must be finite so that `score - max` never produces NaN even when an +/// entire warp has only padded positions. +constexpr float kPadScore = -FLT_MAX; + +/// \brief Online compress 128 prefill. Two passes share this body: +/// - `kWrite=false` (compress pass): handles segments that close a chunk. +/// May load prior partial state from the buffer, but never writes to it, +/// so concurrent blocks can read the same slot without racing. +/// - `kWrite=true` (write pass): handles the trailing partial segment of each +/// batch. Each batch contributes at most one such plan, so concurrent blocks +/// touch disjoint buffer slots. +/// +/// The two passes MUST run as separate kernel launches (in stream order) so +/// that all reads in pass 1 finish before any writes in pass 2 start. +template +__global__ __launch_bounds__(kPrefillBlockSize, 2) // + void flash_c128_online_prefill(const __grid_constant__ Compress128OnlinePrefillParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); + + /// NOTE: the compiler folds the if-else at compile time. + const auto num_plans = kWrite ? params.num_write : params.num_compress; + const auto plan_ptr = kWrite ? params.write_plan : params.compress_plan; + const uint32_t global_id = blockIdx.x; + const uint32_t global_pid = global_id / kNumSplit; // plan id + const uint32_t global_sid = global_id % kNumSplit; // split id + if (global_pid >= num_plans) return; + const auto [ragged_id, batch_id, position, window_len] = plan_ptr[global_pid]; + if (ragged_id == 0xFFFFFFFFu) [[unlikely]] + return; + + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + const int32_t split_offset = global_sid * kTileDim; // int32 is enough + + const auto kv_score_buffer = static_cast(params.kv_score_buffer); + const auto kv_score_input = static_cast(params.kv_score_input); + const auto kv_compressed_output = static_cast(params.kv_compressed_output); + const auto score_bias_base = static_cast(params.score_bias); + + constexpr int64_t kElementSize = kHeadDim * 2; // | kv | score | + const uint32_t chunk_offset = (position % 128u) + 1u - window_len; + const uint32_t window_end = chunk_offset + window_len; // exclusive, in [1, 128] + const int32_t segment_start = ragged_id - (position % 128u); // can be negative, but safe + const int32_t load_index = chunk_offset != 0 ? params.load_indices[batch_id] : -1; + const int32_t store_index = kWrite ? params.indices[batch_id] : -1; + + PDLWaitPrimary(); + + // 2 * 8 = 16 register per elem. in theory we should consume 48 register here + PrefillStorage kv[kElementsPerWarp]; + PrefillStorage score[kElementsPerWarp]; + PrefillStorage bias[kElementsPerWarp]; + const auto warp_offset = warp_id * kElementsPerWarp; + +#pragma unroll + for (uint32_t i = 0; i < kElementsPerWarp; ++i) { + const uint32_t j = i + warp_offset; + if (j >= chunk_offset && j < window_end) { + const auto kv_src_ptr = kv_score_input + (segment_start + j) * kElementSize + split_offset; + const auto score_src_ptr = kv_src_ptr + kHeadDim; + const auto bias_src_ptr = score_bias_base + j * kHeadDim + split_offset; + kv[i].load(kv_src_ptr, lane_id); + score[i].load(score_src_ptr, lane_id); + bias[i].load(bias_src_ptr, lane_id); + } + } + +#pragma unroll + for (uint32_t i = 0; i < kElementsPerWarp; ++i) { + const uint32_t j = i + warp_offset; + const bool is_valid = (j >= chunk_offset && j < window_end); +#pragma unroll + for (uint32_t ii = 0; ii < kTileElements; ++ii) { + score[i][ii] = is_valid ? score[i][ii] + bias[i][ii] : kPadScore; + /// NOTE: must zero out kv on padded slots -- `c128_prefill_forward` + /// computes `kv * exp_score` where `exp_score = expf(-FLT_MAX - max) ??? 0`, + /// and IEEE-754 makes `NaN * 0 = NaN` / `+-inf * 0 = NaN`. An + /// uninitialized register can hold a NaN/inf bit pattern, so without + /// this reset a single padded warp can poison the whole softmax. + kv[i][ii] = is_valid ? kv[i][ii] : 0.0f; + } + } + + __shared__ alignas(16) float seg_kv[kTileDim]; + __shared__ alignas(16) float seg_max[kTileDim]; + __shared__ alignas(16) float seg_sum[kTileDim]; + + c128_prefill_forward(kv, score, seg_kv, seg_max, seg_sum, warp_id, lane_id); + + PDLTriggerSecondary(); + + if (warp_id == 0) { + PrefillStorage out_kv_vec, out_max_vec, out_sum_vec; + out_kv_vec.load(seg_kv, lane_id); + out_max_vec.load(seg_max, lane_id); + out_sum_vec.load(seg_sum, lane_id); + if (chunk_offset != 0) { + /// NOTE: load (max, sum, kv) of the in-progress chunk for this index. + /// `load_indices` may differ from `indices` when the prior partial state + /// lives on a different slot than the slot we ultimately write to. + const auto buf_load = kv_score_buffer + load_index * (kHeadDim * 3) + split_offset; + PrefillStorage buf_max_vec, buf_sum_vec, buf_kv_vec; + buf_max_vec.load(buf_load + 0 * kHeadDim, lane_id); + buf_sum_vec.load(buf_load + 1 * kHeadDim, lane_id); + buf_kv_vec.load(buf_load + 2 * kHeadDim, lane_id); +#pragma unroll + for (uint32_t ii = 0; ii < kTileElements; ++ii) { + const float m1 = buf_max_vec[ii]; + const float s1 = buf_sum_vec[ii]; + const float k1 = buf_kv_vec[ii]; + const float m2 = out_max_vec[ii]; + const float s2 = out_sum_vec[ii]; + const float k2 = out_kv_vec[ii]; + const float new_max = fmaxf(m1, m2); + const float new_s1 = s1 * expf(m1 - new_max); + const float new_s2 = s2 * expf(m2 - new_max); + const float new_sum = new_s1 + new_s2; + const float new_kv = (k1 * new_s1 + k2 * new_s2) / new_sum; + out_max_vec[ii] = new_max; + out_sum_vec[ii] = new_sum; + out_kv_vec[ii] = new_kv; + } + } + + if constexpr (kWrite) { + const auto buf_store = kv_score_buffer + store_index * (kHeadDim * 3) + split_offset; + reinterpret_cast(buf_store + 0 * kHeadDim)[lane_id] = out_max_vec; + reinterpret_cast(buf_store + 1 * kHeadDim)[lane_id] = out_sum_vec; + reinterpret_cast(buf_store + 2 * kHeadDim)[lane_id] = out_kv_vec; + } else { + const auto out_ptr = kv_compressed_output + ragged_id * kHeadDim + split_offset; + reinterpret_cast(out_ptr)[lane_id] = out_kv_vec; + } + } +} + +template +struct FlashCompress128OnlineKernel { + static constexpr auto decode_kernel = flash_c128_online_decode; + template + static constexpr auto prefill_kernel = flash_c128_online_prefill; + static constexpr auto prefill_c_kernel = prefill_kernel; + static constexpr auto prefill_w_kernel = prefill_kernel; + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static constexpr uint32_t kDecodeBlockSize = kHeadDim / 4; + + static void run_decode( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::Optional /* UNUSED */) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) + .with_dtype() + .with_device(device) + .verify(kv_score_buffer); + TensorMatcher({B, kHeadDim * 2}) // kv score input + .with_dtype() + .with_device(device) + .verify(kv_score_input); + TensorMatcher({B, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device) + .verify(ape); + TensorMatcher({B}).with_dtype().with_device(device).verify(indices); + TensorMatcher({B}).with_dtype().with_device(device).verify(seq_lens); + + const auto batch_size = static_cast(B.unwrap()); + const auto params = Compress128OnlineDecodeParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .batch_size = batch_size, + }; + LaunchKernel(batch_size, kDecodeBlockSize, device.unwrap()) // + .enable_pdl(kUsePDL)(decode_kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView compress_plan, + const tvm::ffi::TensorView write_plan, + const tvm::ffi::Optional extra) { + using namespace host; + using host::compress::kOnlinePrefillPlanDim; + using host::compress::OnlinePrefillPlanTensorDtype; + + auto B = SymbolicSize{"batch_size"}; + auto N = SymbolicSize{"num_q_tokens"}; + auto X = SymbolicSize{"compress_tokens"}; + auto Y = SymbolicSize{"write_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) ??? 2D + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({N, kHeadDim * 2}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({N, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + TensorMatcher({B}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + TensorMatcher({X, kOnlinePrefillPlanDim}) // compress plan + .with_dtype() + .with_device(device_) + .verify(compress_plan); + TensorMatcher({Y, kOnlinePrefillPlanDim}) // write plan + .with_dtype() + .with_device(device_) + .verify(write_plan); + + /// NOTE: `extra` is `load_indices`. When the previous partial state lives + /// on a slot different from the destination slot (e.g. paged buffers), the + /// caller must supply this; otherwise it defaults to `indices`. + const auto load_indices = extra.value_or(indices); + TensorMatcher({B}).with_dtype().with_device(device_).verify(load_indices); + + const auto device = device_.unwrap(); + const auto num_c = static_cast(X.unwrap()); + const auto num_w = static_cast(Y.unwrap()); + const auto params = Compress128OnlinePrefillParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .load_indices = static_cast(load_indices.data_ptr()), + .compress_plan = static_cast(compress_plan.data_ptr()), + .write_plan = static_cast(write_plan.data_ptr()), + .num_compress = num_c, + .num_write = num_w, + }; + + /// NOTE: pass 1 reads the buffer (for the first segment of each batch + /// that started mid-chunk) and writes only to `kv_compressed_output`. + /// Pass 2 then writes the trailing partial state of each batch back to + /// the buffer. Stream serialization between the two launches enforces + /// read-before-write on shared buffer slots. + if (const auto num_c_blocks = num_c * kNumSplit) { + LaunchKernel(num_c_blocks, kPrefillBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_c_kernel, params); + } + if (const auto num_w_blocks = num_w * kNumSplit) { + LaunchKernel(num_w_blocks, kPrefillBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_w_kernel, params); + } + } +}; + +} // namespace + +namespace host::compress { + +using OnlinePlanResult = tvm::ffi::Tuple; + +struct OnlinePrefillCompressParams { + OnlinePrefillPlan* __restrict__ compress_plan; + OnlinePrefillPlan* __restrict__ write_plan; + const int64_t* __restrict__ seq_lens; + const int64_t* __restrict__ extend_lens; + uint32_t batch_size; + uint32_t num_tokens; +}; + +/// \brief Build the compress + write plans for online compress 128 prefill. +/// +/// Each batch's `[prefix_len, prefix_len + extend_len)` range is split at +/// 128-aligned boundaries. Every resulting segment falls into one of: +/// - **compress**: closes a 128-chunk (`chunk_offset + window_len == 128`). +/// These plans only read the buffer (when starting mid-chunk) and write the +/// compressed kv to `kv_compressed_output`. +/// - **write**: trailing partial of the batch (`chunk_offset + window_len < 128`). +/// May read the buffer and always writes the new partial state back to it. +/// Each batch produces at most one such plan. +/// +/// The two plans MUST be dispatched as separate kernel launches in stream +/// order so that pass-1 reads of a buffer slot complete before any pass-2 +/// write of the same slot. +inline OnlinePlanResult plan_online_prefill_host(const OnlinePrefillCompressParams& params, const bool use_cuda_graph) { + const auto& [compress_plan, write_plan, seq_lens, extend_lens, batch_size, num_tokens] = params; + + uint32_t counter = 0; + uint32_t compress_count = 0; + uint32_t write_count = 0; + for (const auto i : irange(batch_size)) { + const uint32_t seq_len = static_cast(seq_lens[i]); + const uint32_t extend_len = static_cast(extend_lens[i]); + RuntimeCheck(0 < extend_len && extend_len <= seq_len); + const uint32_t prefix_len = seq_len - extend_len; + const uint32_t end_pos = prefix_len + extend_len; + /// NOTE: split the extend range into per-128-chunk segments. Each segment + /// stays inside one chunk, so the kernel can decide load/store from + /// `chunk_offset` and `window_len` alone. + uint32_t pos = prefix_len; + while (pos < end_pos) { + const uint32_t chunk_start = (pos / 128u) * 128u; + const uint32_t seg_end = std::min(end_pos, chunk_start + 128u); // exclusive + const uint32_t seg_len = seg_end - pos; + const uint32_t chunk_off = pos - chunk_start; + /// NOTE: store last-token coordinates so that downstream consumers + /// (e.g. `fused_norm_rope`) can read `ragged_id` and `position` with the + /// same semantics as `PrefillPlan`. The segment start is recoverable as + /// `ragged_id - window_len + 1` and `position - window_len + 1`. + const uint32_t last_pos = seg_end - 1; + const uint32_t last_ragged = counter + (last_pos - prefix_len); + const auto plan = OnlinePrefillPlan{ + .ragged_id = last_ragged, + .batch_id = i, + .position = last_pos, + .window_len = seg_len, + }; + if (chunk_off + seg_len == 128u) { + // full chunk, must be complete, maybe read the buffer, no write + RuntimeCheck(compress_count < num_tokens); + compress_plan[compress_count++] = plan; + } else { + // last chunk, must be incomplete, maybe read the buffer, must write + RuntimeCheck(write_count < num_tokens); + write_plan[write_count++] = plan; + } + pos = seg_end; + } + counter += extend_len; + } + RuntimeCheck(counter == num_tokens, "input size ", counter, " != num_q_tokens ", num_tokens); + if (!use_cuda_graph) return OnlinePlanResult{compress_count, write_count}; + /// NOTE: pad both plans with sentinel entries so cuda-graph runs always see + /// the same number of blocks. The kernel skips plans whose `ragged_id` is -1. + constexpr auto kInvalid = static_cast(-1); + constexpr auto kInvalidPlan = OnlinePrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; + for (const auto i : irange(compress_count, num_tokens)) { + compress_plan[i] = kInvalidPlan; + } + for (const auto i : irange(write_count, num_tokens)) { + write_plan[i] = kInvalidPlan; + } + return OnlinePlanResult{num_tokens, num_tokens}; +} + +inline OnlinePlanResult plan_online_prefill( + const tvm::ffi::TensorView extend_lens, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView compress_plan, + const tvm::ffi::TensorView write_plan, + const bool use_cuda_graph) { + auto N = SymbolicSize{"batch_size"}; + auto M = SymbolicSize{"num_tokens"}; + auto device = SymbolicDevice{}; + /// NOTE: only host (CPU/cuda-host) planning is implemented for now. The + device.set_options(); + TensorMatcher({N}) // + .with_dtype() + .with_device(device) + .verify(extend_lens) + .verify(seq_lens); + TensorMatcher({M, kOnlinePrefillPlanDim}) // + .with_dtype() + .with_device(device) + .verify(compress_plan) + .verify(write_plan); + const auto params = OnlinePrefillCompressParams{ + .compress_plan = static_cast(compress_plan.data_ptr()), + .write_plan = static_cast(write_plan.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .extend_lens = static_cast(extend_lens.data_ptr()), + .batch_size = static_cast(N.unwrap()), + .num_tokens = static_cast(M.unwrap()), + }; + return plan_online_prefill_host(params, use_cuda_graph); +} + +} // namespace host::compress + +namespace { + +[[maybe_unused]] +constexpr auto& plan_compress_online_prefill = host::compress::plan_online_prefill; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh new file mode 100644 index 0000000000..71e600dc39 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh @@ -0,0 +1,875 @@ +#include +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace { + +using PlanD = device::compress::DecodePlan; +using PlanC = device::compress::CompressPlan; + +// --------------------------------------------------------------------------- +// Decode kernel: 1 token / batch. Each block handles one batch. +// 4 elements per thread -> kBlockSize = head_dim / 4. +// --------------------------------------------------------------------------- + +struct Compress128OnlineDecodeParams { + void* __restrict__ kv_score_buffer; // [num_slots, 1, head_dim * 3] + const void* __restrict__ kv_score_input; // [batch_size, head_dim * 2] + void* __restrict__ kv_compressed_output; // [batch_size, head_dim] + const void* __restrict__ score_bias; // [128, head_dim] + const PlanD* __restrict__ plan_d; + uint32_t batch_size; +}; + +template +__global__ void flash_c128_online_decode_v2(const __grid_constant__ Compress128OnlineDecodeParams params) { + using namespace device; + constexpr uint32_t kVecSize = 4; + constexpr uint32_t kBlockSize = kHeadDim / kVecSize; + using Vec = AlignedVector; + const auto gmem = tile::Memory::cta(kBlockSize); + const auto batch_id = blockIdx.x; + if (batch_id >= params.batch_size) return; + + // Wait for the plan-finalize kernel to publish `plan.read_page_0 / write_loc` + // before reading the plan. The plan kernel runs on the same stream and does + // NOT issue a PDL trigger, so launching this kernel with PDL means our + // pre-wait global reads can race with the plan kernel's writes. + PDLWaitPrimary(); + + const auto plan = params.plan_d[batch_id]; + const auto pos_in_chunk = (plan.seq_len - 1) % 128; + + const auto kv_score_buffer = static_cast(params.kv_score_buffer); + const auto kv_score_input = static_cast(params.kv_score_input); + const auto kv_load_buf = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3); + const auto kv_store_buf = kv_score_buffer + plan.write_loc * (kHeadDim * 3); + const auto kv_src = kv_score_input + batch_id * (kHeadDim * 2); + + // Buffer layout: [max | sum | kv] (slot 0 / 1 / 2 of the head_dim*3 row). + const auto new_kv_vec = gmem.load(kv_src, 0); + const auto new_score_raw_vec = gmem.load(kv_src, 1); + const auto bias_vec = gmem.load(params.score_bias, pos_in_chunk); + + Vec out_kv_vec; + Vec out_max_vec; + Vec out_sum_vec; + if (pos_in_chunk != 0) { + // Mid-chunk: combine prior partial state with the new token. + const auto max_score_vec = gmem.load(kv_load_buf, 0); + const auto sum_score_vec = gmem.load(kv_load_buf, 1); + const auto old_kv_vec = gmem.load(kv_load_buf, 2); +#pragma unroll + for (uint32_t i = 0; i < kVecSize; ++i) { + const auto old_max = max_score_vec[i]; + const auto old_kv = old_kv_vec[i]; + const auto new_score = new_score_raw_vec[i] + bias_vec[i]; + const auto new_kv = new_kv_vec[i]; + const auto new_max = fmaxf(old_max, new_score); + const auto old_sum = sum_score_vec[i] * expf(old_max - new_max); + const auto new_exp = expf(new_score - new_max); + const auto new_sum = old_sum + new_exp; + out_kv_vec[i] = (old_kv * old_sum + new_kv * new_exp) / new_sum; + out_max_vec[i] = new_max; + out_sum_vec[i] = new_sum; + } + } else { + // First token of a new chunk: state == this token alone. +#pragma unroll + for (uint32_t i = 0; i < kVecSize; ++i) { + out_kv_vec[i] = new_kv_vec[i]; + out_max_vec[i] = new_score_raw_vec[i] + bias_vec[i]; + out_sum_vec[i] = 1.0f; + } + } + + if (pos_in_chunk == 127) { + // Chunk just closed: emit compressed kv, no buffer update. + const auto kv_out = static_cast(params.kv_compressed_output) + batch_id * kHeadDim; + gmem.store(kv_out, out_kv_vec); + } else { + gmem.store(kv_store_buf, out_max_vec, 0); + gmem.store(kv_store_buf, out_sum_vec, 1); + gmem.store(kv_store_buf, out_kv_vec, 2); + } +} + +// --------------------------------------------------------------------------- +// Prefill kernel: 1 segment / block. Two passes (compress + write) share the +// kernel template, parameterized by `kWrite`. +// 16 warps per block; each warp handles 8 of the 128 chunk positions. +// --------------------------------------------------------------------------- + +constexpr int32_t kTileElements = 2; // split along head-dim +constexpr int32_t kElementsPerWarp = 8; // split along the 128-chunk +constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; +constexpr uint32_t kPrefillBlockSize = device::kWarpThreads * kNumWarps; +using PrefillStorage = device::AlignedVector; + +struct Compress128OnlinePrefillParams { + void* __restrict__ kv_score_buffer; // [num_slots, 1, head_dim * 3] + const void* __restrict__ kv_score_input; // [num_q_tokens, head_dim * 2] + void* __restrict__ kv_compressed_output; // [num_compress, head_dim] + const void* __restrict__ score_bias; // [128, head_dim] + const PlanC* __restrict__ plan_c; // close-chunk segments + const PlanC* __restrict__ plan_w; // trailing partial segments + uint32_t num_compress; + uint32_t num_write; +}; + +struct Compress128SharedBuffer { + using Storage = device::AlignedVector; + Storage data[kNumWarps][device::kWarpThreads + 1]; // +1 to avoid bank conflict + SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { + return data[warp_id][lane_id]; + } + SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { + return data[warp_id][lane_id][tile_id]; + } +}; + +/// \brief Sentinel score for padded positions in a 128-segment. +constexpr float kPadScore = -FLT_MAX; + +[[maybe_unused]] +SGL_DEVICE void c128_prefill_segment_softmax( + const PrefillStorage (&kv)[kElementsPerWarp], + const PrefillStorage (&score)[kElementsPerWarp], + float* seg_kv, + float* seg_max, + float* seg_sum, + const uint32_t warp_id, + const uint32_t lane_id) { + using namespace device; + + // Per-warp running state (max, sum, kv) for kTileElements head-dim slots. + using TmpStorage = typename Compress128SharedBuffer::Storage; + __shared__ Compress128SharedBuffer s_local_val_max; + __shared__ Compress128SharedBuffer s_local_exp_sum; + __shared__ Compress128SharedBuffer s_local_product; + + TmpStorage tmp_val_max; + TmpStorage tmp_exp_sum; + TmpStorage tmp_product; + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + float score_fp32[kElementsPerWarp]; +#pragma unroll + for (int32_t j = 0; j < kElementsPerWarp; ++j) { + score_fp32[j] = score[j][i]; + } + float max_value = score_fp32[0]; +#pragma unroll + for (int32_t j = 1; j < kElementsPerWarp; ++j) { + max_value = fmaxf(max_value, score_fp32[j]); + } + float sum_exp_value = 0.0f; + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < kElementsPerWarp; ++j) { + const auto exp_score = expf(score_fp32[j] - max_value); + sum_product += kv[j][i] * exp_score; + sum_exp_value += exp_score; + } + tmp_val_max[i] = max_value; + tmp_exp_sum[i] = sum_exp_value; + tmp_product[i] = sum_product; + } + + // Aligned writes (no bank conflict thanks to `+1` padding). + s_local_val_max(warp_id, lane_id) = tmp_val_max; + s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; + s_local_product(warp_id, lane_id) = tmp_product; + + __syncthreads(); + + // Cross-warp reduction. Same recipe as c128_online.cuh: each block-thread + // pair reduces a (tile_id, lane_id) slot using a kNumWarps-wide warp shuffle. + constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; + constexpr uint32_t kIteration = kReductionCount / kPrefillBlockSize; + static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); + +#pragma unroll + for (uint32_t i = 0; i < kIteration; ++i) { + const uint32_t j = i * kPrefillBlockSize + warp_id * kWarpThreads + lane_id; + const uint32_t local_warp_id = j % kNumWarps; + const uint32_t local_elem_id = j / kNumWarps; + const uint32_t local_tile_id = local_elem_id % kTileElements; + const uint32_t local_lane_id = local_elem_id / kTileElements; + const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); + const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); + const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); + const auto global_val_max = warp::reduce_max(local_val_max); + const auto rescale = expf(local_val_max - global_val_max); + const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); + const auto final_scale = rescale / global_exp_sum; + const auto global_product = warp::reduce_sum(local_product * final_scale); + seg_kv[local_elem_id] = global_product; + seg_max[local_elem_id] = global_val_max; + seg_sum[local_elem_id] = global_exp_sum; + } + __syncthreads(); +} + +/// \brief Online compress 128 prefill v2. +/// +/// `kWrite=false` (compress pass): handles segments that close a 128-chunk. +/// Reads optional prior state from `read_page_0` (-1 = none), emits compressed +/// kv to `kv_compressed_output[plan_id]` (compact). +/// `kWrite=true` (write pass) : handles trailing partial segments. +/// Reads optional prior state from `read_page_0` (-1 = none), writes new +/// running state to `read_page_1`. +template +__global__ __launch_bounds__(kPrefillBlockSize, 2) // + void flash_c128_online_prefill_v2(const __grid_constant__ Compress128OnlinePrefillParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static_assert(kHeadDim % kTileDim == 0); + + // Compile-time fold to the right plan list. + const auto num_plans = kWrite ? params.num_write : params.num_compress; + const auto plan_ptr = kWrite ? params.plan_w : params.plan_c; + const uint32_t global_id = blockIdx.x; + const uint32_t global_pid = global_id / kNumSplit; + const uint32_t global_sid = global_id % kNumSplit; + if (global_pid >= num_plans) return; + + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + const int32_t split_offset = global_sid * kTileDim; + + // The previous kernel (plan-finalize stage 1) does NOT issue a PDL trigger, + // so PDLWaitPrimary effectively waits for stage 1 to complete. Read the plan + // AFTER the wait so the freshly-written `read_page_0` (= state-pool slot) is + // visible. Reading it before the wait is a real race -- with PDL enabled the + // kernel can begin executing before stage 1's stores propagate, and we'd see + // the stage-0 batch_id placeholder in `read_page_0` instead of the slot. + PDLWaitPrimary(); + + const auto plan = plan_ptr[global_pid]; + if (plan.is_invalid()) [[unlikely]] + return; + + const auto kv_score_buffer = static_cast(params.kv_score_buffer); + const auto kv_score_input = static_cast(params.kv_score_input); + const auto kv_compressed_output = static_cast(params.kv_compressed_output); + const auto score_bias_base = static_cast(params.score_bias); + + constexpr int64_t kElementSize = kHeadDim * 2; // | kv | score | + + // The plan stores last-token coordinates; segment start is recoverable as + // ragged_id - window_len + 1. + const uint32_t window_len = plan.buffer_len; + const uint32_t position = plan.seq_len - 1; + const uint32_t pos_in_chunk_end = (position % 128u) + 1u; // exclusive, in [1, 128] + const uint32_t chunk_offset = pos_in_chunk_end - window_len; // in [0, 127] + const int32_t segment_start_ragged = static_cast(plan.ragged_id) - static_cast(position % 128u); + + // --- Stage 1: load kv / score / bias for this warp's 8 chunk positions. + PrefillStorage kv[kElementsPerWarp]; + PrefillStorage score[kElementsPerWarp]; + PrefillStorage bias[kElementsPerWarp]; + const uint32_t warp_offset = warp_id * kElementsPerWarp; + +#pragma unroll + for (uint32_t i = 0; i < kElementsPerWarp; ++i) { + const uint32_t j = i + warp_offset; + if (j >= chunk_offset && j < pos_in_chunk_end) { + const auto kv_src_ptr = kv_score_input + (segment_start_ragged + j) * kElementSize + split_offset; + const auto score_src_ptr = kv_src_ptr + kHeadDim; + const auto bias_src_ptr = score_bias_base + j * kHeadDim + split_offset; + kv[i].load(kv_src_ptr, lane_id); + score[i].load(score_src_ptr, lane_id); + bias[i].load(bias_src_ptr, lane_id); + } + } + + // --- Stage 2: pad invalid positions. score = -FLT_MAX, kv = 0 (so that + // kv * exp(score-max) ??? 0 / 0 cleanly without producing NaN/inf). +#pragma unroll + for (uint32_t i = 0; i < kElementsPerWarp; ++i) { + const uint32_t j = i + warp_offset; + const bool is_valid = (j >= chunk_offset && j < pos_in_chunk_end); +#pragma unroll + for (uint32_t ii = 0; ii < kTileElements; ++ii) { + score[i][ii] = is_valid ? score[i][ii] + bias[i][ii] : kPadScore; + kv[i][ii] = is_valid ? kv[i][ii] : 0.0f; + } + } + + // --- Stage 3: warp-tile online softmax over the 128-position chunk. + __shared__ alignas(16) float seg_kv[kTileDim]; + __shared__ alignas(16) float seg_max[kTileDim]; + __shared__ alignas(16) float seg_sum[kTileDim]; + c128_prefill_segment_softmax(kv, score, seg_kv, seg_max, seg_sum, warp_id, lane_id); + + PDLTriggerSecondary(); + + // --- Stage 4: warp 0 folds with prior partial state (if any) and writes. + if (warp_id == 0) { + PrefillStorage out_kv_vec, out_max_vec, out_sum_vec; + out_kv_vec.load(seg_kv, lane_id); + out_max_vec.load(seg_max, lane_id); + out_sum_vec.load(seg_sum, lane_id); + + if (chunk_offset != 0 && plan.read_page_0 >= 0) { + // Combine with prior partial state for this slot. + const auto buf_load = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3) + split_offset; + PrefillStorage buf_max_vec, buf_sum_vec, buf_kv_vec; + buf_max_vec.load(buf_load + 0 * kHeadDim, lane_id); + buf_sum_vec.load(buf_load + 1 * kHeadDim, lane_id); + buf_kv_vec.load(buf_load + 2 * kHeadDim, lane_id); +#pragma unroll + for (uint32_t ii = 0; ii < kTileElements; ++ii) { + const float m1 = buf_max_vec[ii]; + const float s1 = buf_sum_vec[ii]; + const float k1 = buf_kv_vec[ii]; + const float m2 = out_max_vec[ii]; + const float s2 = out_sum_vec[ii]; + const float k2 = out_kv_vec[ii]; + const float new_max = fmaxf(m1, m2); + const float new_s1 = s1 * expf(m1 - new_max); + const float new_s2 = s2 * expf(m2 - new_max); + const float new_sum = new_s1 + new_s2; + const float new_kv = (k1 * new_s1 + k2 * new_s2) / new_sum; + out_max_vec[ii] = new_max; + out_sum_vec[ii] = new_sum; + out_kv_vec[ii] = new_kv; + } + } + + if constexpr (kWrite) { + // For trailing-partial segments the load and store slots collapse to the + // segment's own chunk slot (the request keeps a single in-progress + // chunk's running state at any time), so we reuse `read_page_0`. + const auto buf_store = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3) + split_offset; + reinterpret_cast(buf_store + 0 * kHeadDim)[lane_id] = out_max_vec; + reinterpret_cast(buf_store + 1 * kHeadDim)[lane_id] = out_sum_vec; + reinterpret_cast(buf_store + 2 * kHeadDim)[lane_id] = out_kv_vec; + } else { + // Compact output: one row per compress plan, indexed by `global_pid`. + const auto out_ptr = kv_compressed_output + global_pid * kHeadDim + split_offset; + reinterpret_cast(out_ptr)[lane_id] = out_kv_vec; + } + } +} + +// --------------------------------------------------------------------------- +// Host wrapper: matches the c128_v2 / c4_v2 host API style (run_decode / +// run_prefill methods on a kernel-class template). We only expose `kHeadDim` +// + `kUsePDL`; the dtype is fixed to fp32 for the online state pool. +// --------------------------------------------------------------------------- + +template +struct FlashCompress128OnlineKernel { + static constexpr auto decode_kernel = flash_c128_online_decode_v2; + template + static constexpr auto prefill_kernel = flash_c128_online_prefill_v2; + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static constexpr uint32_t kDecodeBlockSize = kHeadDim / 4; + + static void run_decode( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_d_) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({B, kHeadDim * 2}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({B, kHeadDim}) // kv compressed output (sparse by batch_id) + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + + const auto plan_d = compress::verify_plan_d(plan_d_, B, device_); + const auto batch_size = static_cast(B.unwrap()); + if (batch_size == 0) return; + const auto params = Compress128OnlineDecodeParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_d = plan_d, + .batch_size = batch_size, + }; + LaunchKernel(batch_size, kDecodeBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(decode_kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_c_, + const tvm::ffi::TensorView plan_w_) { + using namespace host; + + auto N = SymbolicSize{"num_q_tokens"}; + auto C = SymbolicSize{"num_c_plans"}; + auto W = SymbolicSize{"num_w_plans"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({N, kHeadDim * 2}) // kv score input (ragged) + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({C, kHeadDim}) // kv compressed output (compact, by plan_c index) + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + + // Both compress and write segments use PlanC layout. plan_c uses + // read_page_1=-1 (unused); plan_w uses read_page_1=store_slot. + const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); + const auto plan_w = compress::verify_plan_c(plan_w_, W, device_); + const auto device = device_.unwrap(); + const auto num_q_tokens = static_cast(N.unwrap()); + const auto num_c = static_cast(C.unwrap()); + const auto num_w = static_cast(W.unwrap()); + RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); + const auto params = Compress128OnlinePrefillParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_c = plan_c, + .plan_w = plan_w, + .num_compress = num_c, + .num_write = num_w, + }; + + // The two passes MUST be serialized in stream order: pass 1 reads slots + // that pass 2 may write to; running them in parallel would race. + if (const auto num_c_blocks = num_c * kNumSplit) { + LaunchKernel(num_c_blocks, kPrefillBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_kernel, params); + } + if (const auto num_w_blocks = num_w * kNumSplit) { + LaunchKernel(num_w_blocks, kPrefillBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_kernel, params); + } + } +}; + +} // namespace + +// =========================================================================== +// Plan builders. Mirrors the offline v2 pattern (`c_plan.cuh`): +// - Decode: a single GPU kernel reads seq_lens / req_to_token / +// req_pool_indices on device and emits the final PlanD tensor in one go. +// - Prefill: stage 0 (host, on CPU pinned memory) splits each batch's +// extend range into per-chunk segments and emits PlanC entries with the +// batch_id stashed in `read_page_0` as a placeholder. Stage 1 is a tiny +// GPU kernel that finalizes `read_page_0` to `req_to_token[rid][chunk_start]`, +// so the slot tensors never leave GPU memory. The online state pool keeps +// a single in-progress chunk per request, so each segment's load and +// store slot collapse to one value (the slot for the segment's own chunk), +// and `read_page_1` is unused. +// =========================================================================== + +namespace host::compress { + +using device::compress::CompressPlan; +using device::compress::DecodePlan; + +// --------------------------------------------------------------------------- +// Decode plan builder. +// --------------------------------------------------------------------------- + +struct OnlineDecodePlanParams { + DecodePlan* __restrict__ plan_d; + const int64_t* __restrict__ seq_lens; + const int64_t* __restrict__ req_pool_indices; + const int32_t* __restrict__ req_to_token; + const int64_t* __restrict__ full_to_swa; // (full_cache_size,) int64 + int64_t stride_r2t; + int32_t swa_page_size; + uint32_t batch_size; +}; + +__global__ void plan_c128_online_decode_kernel(const OnlineDecodePlanParams params) { + const uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= params.batch_size) return; + const auto seq_len = static_cast(params.seq_lens[idx]); + const auto rid = params.req_pool_indices[idx]; + const int32_t chunk_start = static_cast((seq_len - 1u) / 128u * 128u); + const int32_t full_loc = params.req_to_token[rid * params.stride_r2t + chunk_start]; + const int32_t swa_loc = static_cast(params.full_to_swa[full_loc]); + const int32_t slot = swa_loc / params.swa_page_size; + params.plan_d[idx] = DecodePlan{ + .seq_len = seq_len, + .write_loc = slot, + .read_page_0 = slot, + .read_page_1 = -1, + }; +} + +/// \brief Build the decode plan tensor. Caller (Python) pre-allocates +/// `plan_d_dev` as a `(batch_size, 16)` device uint8 tensor; this routine +/// only fills it. See `plan_online_prefill` for the rationale (avoid +/// `ffi::empty` + dlpack roundtrip / PyTorch caching-allocator stream +/// tracking issue that surfaces as IMA in unrelated downstream kernels). +inline void plan_online_decode( + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView req_pool_indices, + const tvm::ffi::TensorView req_to_token, + const tvm::ffi::TensorView full_to_swa, + const tvm::ffi::TensorView plan_d_dev_, + const int32_t swa_page_size) { + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + auto seq_dtype = SymbolicDType{}; + TensorMatcher({B}) // + .with_dtype(seq_dtype) + .with_device(device_) + .verify(seq_lens); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({-1, -1}) // + .with_dtype() + .with_device(device_) + .verify(req_to_token); + TensorMatcher({-1}) // + .with_dtype() + .with_device(device_) + .verify(full_to_swa); + TensorMatcher({B, sizeof(DecodePlan)}) // + .with_dtype() + .with_device(device_) + .verify(plan_d_dev_); + RuntimeCheck(swa_page_size > 0); + + const auto batch_size = static_cast(B.unwrap()); + if (batch_size == 0) return; + + const auto device = device_.unwrap(); + constexpr uint32_t kBlockSize = 256; + const uint32_t num_blocks = host::div_ceil(batch_size, kBlockSize); + const auto stride_r2t = req_to_token.stride(0); + const auto params = OnlineDecodePlanParams{ + .plan_d = static_cast(plan_d_dev_.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .req_pool_indices = static_cast(req_pool_indices.data_ptr()), + .req_to_token = static_cast(req_to_token.data_ptr()), + .full_to_swa = static_cast(full_to_swa.data_ptr()), + .stride_r2t = stride_r2t, + .swa_page_size = swa_page_size, + .batch_size = batch_size, + }; + LaunchKernel(num_blocks, kBlockSize, device)(plan_c128_online_decode_kernel, params); +} + +// --------------------------------------------------------------------------- +// Prefill plan builder: host stage 0 + GPU stage 1. +// --------------------------------------------------------------------------- + +struct OnlinePrefillStage0Params { + CompressPlan* __restrict__ plan_c; + CompressPlan* __restrict__ plan_w; + const int64_t* __restrict__ seq_lens; + const int64_t* __restrict__ extend_lens; + uint32_t batch_size; + uint32_t num_q_tokens; +}; + +inline std::tuple _plan_prefill_partial(const OnlinePrefillStage0Params& p) { + uint32_t counter = 0; + uint32_t compress_count = 0; + uint32_t write_count = 0; + for (const auto i : irange(p.batch_size)) { + const uint32_t seq_len = static_cast(p.seq_lens[i]); + const uint32_t extend_len = static_cast(p.extend_lens[i]); + RuntimeCheck(0 < extend_len && extend_len <= seq_len); + const uint32_t prefix_len = seq_len - extend_len; + const uint32_t end_pos = prefix_len + extend_len; + + uint32_t pos = prefix_len; + while (pos < end_pos) { + const uint32_t chunk_start = (pos / 128u) * 128u; + const uint32_t seg_end = std::min(end_pos, chunk_start + 128u); // exclusive + const uint32_t seg_len = seg_end - pos; + const uint32_t chunk_off = pos - chunk_start; + const uint32_t last_pos = seg_end - 1; + const uint32_t last_ragged = counter + (last_pos - prefix_len); + RuntimeCheck(last_ragged < (1u << 16), "PlanC.ragged_id is uint16; ragged ", last_ragged, " overflows"); + RuntimeCheck(seg_len <= 128u); + // Stash batch_id in `read_page_0` for stage 1 to translate. A + // chunk-aligned segment never loads, so we still need stage 1 to fill + // a slot in -- the kernel keys the load on `chunk_offset != 0`. + const auto plan = CompressPlan{ + .seq_len = last_pos + 1u, + .ragged_id = static_cast(last_ragged), + .buffer_len = static_cast(seg_len), + .read_page_0 = static_cast(i), // batch_id placeholder + .read_page_1 = -1, // unused, kept so MSB layout is stable + }; + if (chunk_off + seg_len == 128u) { + // close-chunk segment + RuntimeCheck(compress_count < p.num_q_tokens); + p.plan_c[compress_count++] = plan; + } else { + // trailing partial segment + RuntimeCheck(write_count < p.num_q_tokens); + p.plan_w[write_count++] = plan; + } + pos = seg_end; + } + counter += extend_len; + } + RuntimeCheck(counter == p.num_q_tokens, "input size ", counter, " != num_q_tokens ", p.num_q_tokens); + return std::tuple{compress_count, write_count}; +} + +struct OnlinePrefillStage1Params { + CompressPlan* __restrict__ plan_c; + CompressPlan* __restrict__ plan_w; + const int64_t* __restrict__ req_pool_indices; // (batch_size,) + const int32_t* __restrict__ req_to_token; // (num_reqs, max_tokens) + const int64_t* __restrict__ full_to_swa; // (full_cache_size,) + int64_t stride_r2t; + int32_t swa_page_size; + uint32_t num_c; + uint32_t num_w; +}; + +__global__ void plan_c128_online_prefill_kernel(const OnlinePrefillStage1Params params) { + const uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t total = params.num_c + params.num_w; + if (idx >= total) return; + + const bool is_compress = idx < params.num_c; + CompressPlan* const plan_ptr = is_compress ? ¶ms.plan_c[idx] : ¶ms.plan_w[idx - params.num_c]; + auto plan = *plan_ptr; + const auto batch_id = plan.read_page_0; + const auto rid = params.req_pool_indices[batch_id]; + const int32_t position = static_cast(plan.seq_len - 1u); + const int32_t chunk_start = (position / 128) * 128; + const int32_t full_loc = params.req_to_token[rid * params.stride_r2t + chunk_start]; + const int32_t swa_loc = static_cast(params.full_to_swa[full_loc]); + plan.read_page_0 = swa_loc / params.swa_page_size; + *plan_ptr = plan; +} + +using OnlinePrefillPlan = tvm::ffi::Tuple; + +inline OnlinePrefillPlan plan_online_prefill( + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView extend_lens, + const tvm::ffi::TensorView req_pool_indices, + const tvm::ffi::TensorView req_to_token, + const tvm::ffi::TensorView full_to_swa, + const tvm::ffi::TensorView plan_c_pin, + const tvm::ffi::TensorView plan_w_pin, + const tvm::ffi::TensorView plan_c_dev_, + const tvm::ffi::TensorView plan_w_dev_, + const int32_t swa_page_size) { + auto B = SymbolicSize{"batch_size"}; + auto N = SymbolicSize{"num_q_tokens"}; + auto cpu = SymbolicDevice{}; + auto device_ = SymbolicDevice{}; + cpu.set_options(); + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(cpu) + .verify(seq_lens) + .verify(extend_lens); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({-1, -1}) // + .with_dtype() + .with_device(device_) + .verify(req_to_token); + TensorMatcher({-1}) // + .with_dtype() + .with_device(device_) + .verify(full_to_swa); + TensorMatcher({N, sizeof(CompressPlan)}) // + .with_dtype() + .with_device(cpu) + .verify(plan_c_pin) + .verify(plan_w_pin); + TensorMatcher({N, sizeof(CompressPlan)}) // + .with_dtype() + .with_device(device_) + .verify(plan_c_dev_) + .verify(plan_w_dev_); + + const auto stage0_params = OnlinePrefillStage0Params{ + .plan_c = static_cast(plan_c_pin.data_ptr()), + .plan_w = static_cast(plan_w_pin.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .extend_lens = static_cast(extend_lens.data_ptr()), + .batch_size = static_cast(B.unwrap()), + .num_q_tokens = static_cast(N.unwrap()), + }; + + // Debug instrumentation: SGLANG_DEBUG_C128_ONLINE_GUARD=1 wraps stage 0 + // with redzone + post-write magic-check on the pin buffers, plus a strict + // upper-bound check on `batch_size` and `num_q_tokens`. If stage 0 has a + // CPU OOB this trips a clear panic at the offending byte instead of a + // delayed CUDA IMA from corrupted heap memory. + static const bool kGuard = []() { + const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_GUARD"); + return v != nullptr && v[0] == '1'; + }(); + if (kGuard) { + RuntimeCheck(stage0_params.batch_size <= 65536u, "batch_size out of bound: ", stage0_params.batch_size); + RuntimeCheck(stage0_params.num_q_tokens <= 65536u, "num_q_tokens out of bound: ", stage0_params.num_q_tokens); + // Stamp the pin buffers with 0xAB so we can detect any byte still 0xAB + // beyond what stage 0 should have written (= OOB never reached, that's fine) + // or any byte BEYOND num_q_tokens*16 written to (= true OOB into + // adjacent allocation). + auto* pc = static_cast(plan_c_pin.data_ptr()); + auto* pw = static_cast(plan_w_pin.data_ptr()); + const auto bytes = static_cast(N.unwrap()) * sizeof(CompressPlan); + std::memset(pc, 0xAB, bytes); + std::memset(pw, 0xAB, bytes); + } + + const auto [num_c, num_w] = _plan_prefill_partial(stage0_params); + + if (kGuard) { + // Verify stage 0 wrote ONLY to the [0, num_c*16) and [0, num_w*16) prefix. + auto* pc = static_cast(plan_c_pin.data_ptr()); + auto* pw = static_cast(plan_w_pin.data_ptr()); + const auto end_c = static_cast(num_c) * sizeof(CompressPlan); + const auto end_w = static_cast(num_w) * sizeof(CompressPlan); + const auto pin_bytes = static_cast(N.unwrap()) * sizeof(CompressPlan); + for (size_t k = end_c; k < pin_bytes; ++k) { + RuntimeCheck( + pc[k] == 0xAB, + "GUARD: plan_c_pin OOB write at byte ", + k, + " (num_c=", + num_c, + ", num_q_tokens=", + N.unwrap(), + ")"); + } + for (size_t k = end_w; k < pin_bytes; ++k) { + RuntimeCheck( + pw[k] == 0xAB, + "GUARD: plan_w_pin OOB write at byte ", + k, + " (num_w=", + num_w, + ", num_q_tokens=", + N.unwrap(), + ")"); + } + } + + const auto device = device_.unwrap(); + // Out-params pre-allocated by Python. Cast to typed pointers for use. + auto* const plan_c_dev_ptr = static_cast(plan_c_dev_.data_ptr()); + auto* const plan_w_dev_ptr = static_cast(plan_w_dev_.data_ptr()); + + if (const auto total = num_c + num_w) { + const auto stream = LaunchKernel::resolve_device(device); + // SGLANG_DEBUG_C128_ONLINE_SYNC_H2D=1 forces a synchronous H2D copy. + static const bool kSyncH2D = []() { + const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_SYNC_H2D"); + return v != nullptr && v[0] == '1'; + }(); + // SGLANG_DEBUG_C128_ONLINE_NO_H2D=1 skips the H2D copy entirely (debug only). + static const bool kNoH2D = []() { + const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_NO_H2D"); + return v != nullptr && v[0] == '1'; + }(); + const auto copy_to_device = [stream](void* dst, void* src, int64_t count) { + if (kNoH2D) return; + const auto bytes = count * sizeof(CompressPlan); + if (kSyncH2D) { + RuntimeDeviceCheck(::cudaMemcpy(dst, src, bytes, ::cudaMemcpyHostToDevice)); + } else { + RuntimeDeviceCheck(::cudaMemcpyAsync(dst, src, bytes, ::cudaMemcpyHostToDevice, stream)); + } + }; + if (num_c) copy_to_device(plan_c_dev_ptr, plan_c_pin.data_ptr(), num_c); + if (num_w) copy_to_device(plan_w_dev_ptr, plan_w_pin.data_ptr(), num_w); + + const auto stage1_params = OnlinePrefillStage1Params{ + .plan_c = plan_c_dev_ptr, + .plan_w = plan_w_dev_ptr, + .req_pool_indices = static_cast(req_pool_indices.data_ptr()), + .req_to_token = static_cast(req_to_token.data_ptr()), + .full_to_swa = static_cast(full_to_swa.data_ptr()), + .stride_r2t = req_to_token.stride(0), + .swa_page_size = swa_page_size, + .num_c = num_c, + .num_w = num_w, + }; + constexpr uint32_t kBlockSize = 128; + const auto num_blocks = host::div_ceil(total, kBlockSize); + LaunchKernel(num_blocks, kBlockSize, device)(plan_c128_online_prefill_kernel, stage1_params); + } + return OnlinePrefillPlan{num_c, num_w}; +} + +} // namespace host::compress + +namespace { + +[[maybe_unused]] +constexpr auto& plan_compress_128_online_decode = host::compress::plan_online_decode; +[[maybe_unused]] +constexpr auto& plan_compress_128_online_prefill = host::compress::plan_online_prefill; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh new file mode 100644 index 0000000000..31353e6a15 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh @@ -0,0 +1,448 @@ +/** + * \brief Here's some dimension info for the main buffer used in C128 prefill and decode. + * + * kv_buffer: [num_indices, 128, head_dim * 2] + * - last dimension layout: | kv | score | + * kv_input: [batch_size, head_dim * 2] + * kv_output: [batch_size, head_dim] + * score_bias (ape): [128, head_dim] + * plan_c/plan_w: [variable length] + * + * For prefill, batch_size = num_q_tokens + */ + +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include + +namespace { + +using PlanD = device::compress::DecodePlan; +using PlanC = device::compress::CompressPlan; +using PlanW = device::compress::WritePlan; + +/// \brief Each thread will handle this many elements (split along head_dim) +constexpr int32_t kTileElements = 2; +/// \brief Each warp will handle this many elements (split along 128) +constexpr int32_t kElementsPerWarp = 8; +constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; +constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps; +constexpr uint32_t kWriteBlockSize = 128; // one warp per write + +/// \brief Need to reduce register usage to increase occupancy +#define C128_KERNEL __global__ __launch_bounds__(kBlockSize, 2) +#define WRITE_KERNEL __global__ __launch_bounds__(kWriteBlockSize, 16) + +struct Compress128DecodeParams { + void* __restrict__ kv_buffer; + const void* __restrict__ kv_input; + void* __restrict__ kv_output; + const void* __restrict__ score_bias; + const PlanD* __restrict__ plan_d; + uint32_t batch_size; +}; + +struct Compress128PrefillParams { + void* __restrict__ kv_buffer; + const void* __restrict__ kv_input; + void* __restrict__ kv_output; + const void* __restrict__ score_bias; + const PlanC* __restrict__ plan_c; + const PlanW* __restrict__ plan_w; + uint32_t num_compress; + uint32_t num_write; +}; + +struct Compress128SharedBuffer { + using Storage = device::AlignedVector; + Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict + SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { + return data[warp_id][lane_id]; + } + SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { + return data[warp_id][lane_id][tile_id]; + } +}; + +template +struct C128Trait { + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 + static constexpr int64_t kHeadDim = kHeadDim_; + static constexpr int64_t kScoreOffset = kHeadDim; + static constexpr int64_t kElementSize = kHeadDim * 2; + static constexpr int64_t kPageElementSize = 128 * kElementSize; // page size = 128 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static_assert(kHeadDim % kTileDim == 0); +}; + +template +SGL_DEVICE void c128_forward( + const InFloat* kv_buf, // [128n, 128n + 127] + const InFloat* kv_src, // ragged pointer at position = 128n + 127 + OutFloat* kv_out, + const InFloat* score_bias, + const int32_t buffer_len) { + using namespace device; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + + /// NOTE: part 1: load kv + score + using StorageIn = AlignedVector; + const auto gmem_in = tile::Memory{lane_id, kWarpThreads}; + StorageIn kv[kElementsPerWarp]; + StorageIn score[kElementsPerWarp]; + StorageIn bias[kElementsPerWarp]; + const int32_t warp_offset = warp_id * kElementsPerWarp; + +#pragma unroll + for (int32_t i = 0; i < 8; ++i) { + const int32_t j = i + warp_offset; + bias[i] = gmem_in.load(score_bias + j * Trait::kHeadDim); + } + + const auto kv_start = kv_src - 127 * Trait::kElementSize; // point to start + +#pragma unroll + for (int32_t i = 0; i < kElementsPerWarp; ++i) { + const int32_t j = i + warp_offset; + __builtin_assume(j < 128); + const auto src = j < buffer_len ? kv_buf : kv_start; + kv[i] = gmem_in.load(src + j * Trait::kElementSize); + score[i] = gmem_in.load(src + j * Trait::kElementSize + Trait::kScoreOffset); + } + + /// NOTE: part 2: safe online softmax + weighted sum + using TmpStorage = typename Compress128SharedBuffer::Storage; + __shared__ Compress128SharedBuffer s_local_val_max; + __shared__ Compress128SharedBuffer s_local_exp_sum; + __shared__ Compress128SharedBuffer s_local_product; + + TmpStorage tmp_val_max; + TmpStorage tmp_exp_sum; + TmpStorage tmp_product; + + float score_fp32[kTileElements][kElementsPerWarp]; + + // convert to fp32 and apply bias first +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + for (int32_t j = 0; j < kElementsPerWarp; ++j) { + score_fp32[i][j] = cast(score[j][i]) + cast(bias[j][i]); + } + } + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + const auto& score = score_fp32[i]; + float max_value = score[0]; + float sum_exp_value = 0.0f; + +#pragma unroll + for (int32_t j = 1; j < kElementsPerWarp; ++j) { + const auto fp32_score = score[j]; + max_value = fmaxf(max_value, fp32_score); + } + + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + const auto fp32_score = score[j]; + const auto exp_score = expf(fp32_score - max_value); + sum_product += cast(kv[j][i]) * exp_score; + sum_exp_value += exp_score; + } + + tmp_val_max[i] = max_value; + tmp_exp_sum[i] = sum_exp_value; + tmp_product[i] = sum_product; + } + + // naturally aligned, so no bank conflict + s_local_val_max(warp_id, lane_id) = tmp_val_max; + s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; + s_local_product(warp_id, lane_id) = tmp_product; + + __syncthreads(); + + /// NOTE: part 3: online softmax + /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce + /// each reduce will consume `kNumWarps` threads (use partial warp reduction) + constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; + constexpr uint32_t kIteration = kReductionCount / kBlockSize; + + PDLTriggerSecondary(); + +#pragma unroll + for (uint32_t i = 0; i < kIteration; ++i) { + /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` + const uint32_t j = i * kBlockSize + warp_id * kWarpThreads + lane_id; + /// NOTE: Range `[0, kNumWarps)` + const uint32_t local_warp_id = j % kNumWarps; + /// NOTE: Range `[0, kTileElements * kWarpThreads)` + const uint32_t local_elem_id = j / kNumWarps; + /// NOTE: Range `[0, kTileElements)` + const uint32_t local_tile_id = local_elem_id % kTileElements; + /// NOTE: Range `[0, kWarpThreads)` + const uint32_t local_lane_id = local_elem_id / kTileElements; + /// NOTE: each warp will access the whole tile (all `kTileElements`) + /// and for different lanes, the memory access only differ in `local_warp_id` + /// so there's no bank conflict in shared memory access. + static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); + const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); + const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); + const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); + const auto global_val_max = warp::reduce_max(local_val_max); + const auto rescale = expf(local_val_max - global_val_max); + const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); + const auto final_scale = rescale / global_exp_sum; + const auto global_product = warp::reduce_sum(local_product * final_scale); + kv_out[local_elem_id] = cast(global_product); + } +} + +template +SGL_DEVICE void c128_write_decode(InFloat* kv_buf, const InFloat* kv_src) { + using namespace device; + + using Storage = AlignedVector; + const auto gmem = tile::Memory::warp(); + + Storage data[2]; +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + data[i] = gmem.load(kv_src + Trait::kHeadDim * i); + } +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + gmem.store(kv_buf + Trait::kHeadDim * i, data[i]); + } +} + +template +C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) { + using namespace device; + using Trait = C128Trait; + + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t global_bid = blockIdx.x / Trait::kNumSplit; // batch id + const uint32_t global_sid = blockIdx.x % Trait::kNumSplit; // split id + const int64_t split_offset = global_sid * Trait::kTileDim; + if (global_bid >= params.batch_size) return; + + const auto plan = params.plan_d[global_bid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_output = static_cast(params.kv_output) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + const auto score_bias = static_cast(params.score_bias) + split_offset; + + const auto kv_src = kv_input + global_bid * Trait::kElementSize; + const auto kv_out = kv_output + global_bid * Trait::kHeadDim; + const auto kv_buf = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; + const auto kv_dst = kv_buffer + plan.write_loc * Trait::kElementSize; + + PDLWaitPrimary(); + // the write warp must match the load warp in the following `c128_forward` + if (warp_id == kNumWarps - 1) { + c128_write_decode(kv_dst, kv_src); + } + if (plan.write_loc % 128 == 127) { + c128_forward(kv_buf, kv_src, kv_out, score_bias, 128); + } +} + +// compress kernel +template +C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { + using namespace device; + using Trait = C128Trait; + + const uint32_t global_pid = blockIdx.x / Trait::kNumSplit; // plan id + const uint32_t global_sid = blockIdx.x % Trait::kNumSplit; // split id + const int64_t split_offset = global_sid * Trait::kTileDim; + if (global_pid >= params.num_compress) return; + + const auto plan = params.plan_c[global_pid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_output = static_cast(params.kv_output) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + const auto score_bias = static_cast(params.score_bias) + split_offset; + if (plan.is_invalid()) return; + + const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; + // Compact output: one row per compress plan, indexed by `global_pid`. + const auto kv_out = kv_output + global_pid * Trait::kHeadDim; + const auto kv_buf = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; + PDLWaitPrimary(); + c128_forward(kv_buf, kv_src, kv_out, score_bias, plan.buffer_len); +} + +template +WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { + using namespace device; + using Trait = C128Trait; + using StorageIn = AlignedVector; + + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id + const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id + // split the contiguous `kHeadDim * 2` into `kNumSplit` tiles + // each warp handles 1 contiguous tile (in contrast, decode handle the strided head_dim) + const int64_t split_offset = global_sid * (Trait::kTileDim * 2); + if (global_pid >= params.num_write) return; + + const auto plan = params.plan_w[global_pid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + if (plan.is_invalid()) return; + + // each warp will handle a contiguous region + const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; + const auto kv_buf = kv_buffer + plan.write_loc * Trait::kElementSize; + const auto gmem = tile::Memory::warp(); + + PDLWaitPrimary(); + StorageIn data[2]; +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + data[i] = gmem.load(kv_src, i); + } + PDLTriggerSecondary(); +#pragma unroll + for (int32_t i = 0; i < 2; ++i) { + gmem.store(kv_buf, data[i], i); + } +} + +template +struct FlashCompress128Kernel { + static constexpr auto decode_kernel = flash_c128_decode; + static constexpr auto prefill_c_kernel = flash_c128_prefill; + static constexpr auto prefill_w_kernel = write_c128_prefill; + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + using Trait = C128Trait; + + static void run_decode( + const tvm::ffi::TensorView kv_buffer, + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_d_) { + using namespace host; + + auto N = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 128, Trait::kElementSize}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_buffer); + TensorMatcher({N, Trait::kElementSize}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_input); + TensorMatcher({N, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + + const auto plan_d = compress::verify_plan_d(plan_d_, N, device_); + const auto batch_size = static_cast(N.unwrap()); + const auto params = Compress128DecodeParams{ + .kv_buffer = kv_buffer.data_ptr(), + .kv_input = kv_input.data_ptr(), + .kv_output = kv_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_d = plan_d, + .batch_size = batch_size, + }; + const uint32_t num_blocks = batch_size * kNumSplit; + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(decode_kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_buffer, + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_c_, + const tvm::ffi::TensorView plan_w_) { + using namespace host; + + auto N = SymbolicSize{"num_q_tokens"}; + auto C = SymbolicSize{"num_c_plans"}; + auto W = SymbolicSize{"num_w_plans"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 128, Trait::kElementSize}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_buffer); + TensorMatcher({N, Trait::kElementSize}) // kv score input (ragged) + .with_dtype() + .with_device(device_) + .verify(kv_input); + TensorMatcher({C, kHeadDim}) // kv compressed output (compact) + .with_dtype() + .with_device(device_) + .verify(kv_output); + TensorMatcher({128, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + + const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); + const auto plan_w = compress::verify_plan_w(plan_w_, W, device_); + const auto device = device_.unwrap(); + const auto num_q_tokens = static_cast(N.unwrap()); + const auto num_c = static_cast(C.unwrap()); + const auto num_w = static_cast(W.unwrap()); + const auto params = Compress128PrefillParams{ + .kv_buffer = kv_buffer.data_ptr(), + .kv_input = kv_input.data_ptr(), + .kv_output = kv_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_c = plan_c, + .plan_w = plan_w, + .num_compress = num_c, + .num_write = num_w, + }; + RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); + if (const auto num_c_blocks = num_c * kNumSplit) { + constexpr auto kBlockSize_C = kBlockSize; + LaunchKernel(num_c_blocks, kBlockSize_C, device) // + .enable_pdl(kUsePDL)(prefill_c_kernel, params); + } + constexpr uint32_t kWarpsPerWriteBlock = kWriteBlockSize / device::kWarpThreads; + if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerWriteBlock)) { + constexpr auto kBlockSize_W = kWriteBlockSize; + LaunchKernel(num_w_blocks, kBlockSize_W, device) // + .enable_pdl(kUsePDL)(prefill_w_kernel, params); + } + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh new file mode 100644 index 0000000000..145ab1fb08 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh @@ -0,0 +1,549 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include + +namespace { + +using Plan4 = device::compress::PrefillPlan; +using IndiceT = int32_t; + +/// \brief Each thread will handle this many elements (split along head_dim) +constexpr int kTileElements = 4; + +/// \brief Need to improve register usage to reduce latency +#define C4_KERNEL __global__ __launch_bounds__(128, 4) + +enum class PageMode { + RingBuffer = 8, + Page4Align = 4, +}; + +struct alignas(16) C4IndexBundle { + int32_t load_first_page; + int32_t load_second_page; + int32_t write_first_page; + int32_t last_position; +}; + +struct Compress4DecodeParams { + /** + * \brief Shape: `[num_indices, 8, head_dim * 4]` \n + * last dimension layout: + * | kv overlap | kv | score overlap | score | + */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[batch_size, head_dim * 4]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[batch_size, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[8, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, ]` */ + const IndiceT* __restrict__ seq_lens; + /** \brief Shape: `[batch_size, 1]` */ + const int32_t* __restrict__ extra; + /** \NOTE: `batch_size` <= `num_indices` */ + uint32_t batch_size; +}; + +struct Compress4PrefillParams { + /** + * \brief Shape: `[num_indices, 8, head_dim * 4]` \n + * last dimension layout: + * | kv overlap | kv | score overlap | score | + */ + void* __restrict__ kv_score_buffer; + /** \brief Shape: `[num_q_tokens, head_dim * 4]` */ + const void* __restrict__ kv_score_input; + /** \brief Shape: `[num_q_tokens, head_dim]` */ + void* __restrict__ kv_compressed_output; + /** \brief Shape: `[8, head_dim]` (called `ape`) */ + const void* __restrict__ score_bias; + /** \brief Shape: `[batch_size, ]`*/ + const IndiceT* __restrict__ indices; + /** \brief Shape: `[batch_size, 4]` */ + const C4IndexBundle* __restrict__ extra; + /** \brief The following part is plan info. */ + + const Plan4* __restrict__ compress_plan; + const Plan4* __restrict__ write_plan; + uint32_t num_compress; + uint32_t num_write; +}; + +template +SGL_DEVICE void c4_write( + T* kv_score_buf, // + const T* kv_score_src, + const int64_t head_dim, + const int32_t write_pos) { + using namespace device; + + using Storage = AlignedVector; + const auto element_size = head_dim * 4; + const auto gmem = tile::Memory::warp(); + kv_score_buf += write_pos * element_size; + + /// NOTE: Layout | [0] = kv overlap | [1] = kv | [2] = score overlap | [3] = score | + Storage kv_score[4]; +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + kv_score[i] = gmem.load(kv_score_src + head_dim * i); + } +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + gmem.store(kv_score_buf + head_dim * i, kv_score[i]); + } +} + +template +SGL_DEVICE void c4_forward( + const InFloat* kv_score_buf, + const InFloat* kv_score_src, + OutFloat* kv_out, + const InFloat* score_bias, + const int64_t head_dim, + const int32_t seq_len, + const int32_t window_len, + [[maybe_unused]] const InFloat* kv_score_overlap_buf = nullptr) { + using namespace device; + + const auto element_size = head_dim * 4; + const auto score_offset = head_dim * 2; + const auto overlap_stride = head_dim; + + /// NOTE: part 1: load kv + score + using StorageIn = AlignedVector; + const auto gmem_in = tile::Memory::warp(); + StorageIn kv[8]; + StorageIn score[8]; + StorageIn bias[8]; + +#pragma unroll + for (int32_t i = 0; i < 8; ++i) { + bias[i] = gmem_in.load(score_bias + i * head_dim); + } + +#pragma unroll + for (int32_t i = 0; i < 8; ++i) { + const bool is_overlap = i < 4; + const InFloat* src; + if (i < window_len) { + /// NOTE: `seq_len` must be a multiple of 4 here + if constexpr (kPaged) { + const auto kv_score_ptr = is_overlap ? kv_score_overlap_buf : kv_score_buf; + const int32_t k = i % 4; + src = kv_score_ptr + k * element_size; + } else { + const int32_t k = (seq_len + i) % 8; + src = kv_score_buf + k * element_size; + } + } else { + /// NOTE: k in [-7, 0]. We'll load from the ragged `kv_score_src` + const int32_t k = i - 7; + src = kv_score_src + k * element_size; + } + src += (is_overlap ? 0 : overlap_stride); + kv[i] = gmem_in.load(src); + score[i] = gmem_in.load(src + score_offset); + } + + if (seq_len == 4) { + [[unlikely]]; + constexpr float kFloatNegInf = -1e9f; +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + kv[i].fill(cast(0.0f)); + score[i].fill(cast(kFloatNegInf)); + } + } + + /// NOTE: part 2: safe online softmax + weighted sum + using StorageOut = AlignedVector; + const auto gmem_out = tile::Memory::warp(); + StorageOut result; + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + float score_fp32[8]; + +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + score_fp32[j] = cast(score[j][i]) + cast(bias[j][i]); + } + + float max_value = score_fp32[0]; + float sum_exp_value = 0.0f; + +#pragma unroll + for (int32_t j = 1; j < 8; ++j) { + const auto fp32_score = score_fp32[j]; + max_value = fmaxf(max_value, fp32_score); + } + + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + const auto fp32_score = score_fp32[j]; + const auto exp_score = expf(fp32_score - max_value); + sum_product += cast(kv[j][i]) * exp_score; + sum_exp_value += exp_score; + } + + result[i] = cast(sum_product / sum_exp_value); + } + + gmem_out.store(kv_out, result); +} + +template +C4_KERNEL void flash_c4_decode(const __grid_constant__ Compress4DecodeParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 128 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + constexpr int64_t kElementSize = kHeadDim * 4; // `* 4` due to overlap transform + score + static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); + + const auto& [ + _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score + indices, seq_lens, extra, batch_size // decode info + ] = params; + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_bid = global_wid / kNumSplit; // batch id + const uint32_t global_sid = global_wid % kNumSplit; // split id + + if (global_bid >= batch_size) return; + + const int32_t index = indices[global_bid]; + const int32_t seq_len = seq_lens[global_bid]; + const int64_t split_offset = global_sid * kTileDim; + + // kv score + const auto kv_score_buffer = static_cast(_kv_score_buffer); + + // kv input + const auto kv_score_input = static_cast(_kv_score_input); + const auto kv_src = kv_score_input + global_bid * kElementSize + split_offset; + + // kv output + const auto kv_compressed_output = static_cast(_kv_compressed_output); + const auto kv_out = kv_compressed_output + global_bid * kHeadDim + split_offset; + + // score bias (ape) + const auto score_bias = static_cast(_score_bias) + split_offset; + + PDLWaitPrimary(); + + /// NOTE: `position` = `seq_len - 1`. To avoid underflow, we use `seq_len + page_size - 1` + if constexpr (kMode == PageMode::Page4Align) { + const auto index_prev = extra[global_bid]; + const auto kv_buf = kv_score_buffer + index * (kElementSize * 4) + split_offset; + c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 3) % 4); + if (seq_len % 4 == 0) { + const auto kv_overlap = kv_buf + (index_prev - index) * (kElementSize * 4); + c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, 8, kv_overlap); + } + } else { + static_assert(kMode == PageMode::RingBuffer, "Unsupported PageMode"); + const auto kv_buf = kv_score_buffer + index * (kElementSize * 8) + split_offset; + c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 7) % 8); + if (seq_len % 4 == 0) { + c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, /*window_size=*/8); + } + } + + PDLTriggerSecondary(); +} + +template +C4_KERNEL void flash_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { + using namespace device; + + constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 128 + constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + constexpr int64_t kElementSize = kHeadDim * 4; // `* 4` due to overlap transform + score + static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); + + const auto& [ + _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score + indices, extra, compress_plan, write_plan, num_compress, num_write // prefill plan + ] = params; + + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_pid = global_wid / kNumSplit; // plan id + const uint32_t global_sid = global_wid % kNumSplit; // split id + + /// NOTE: compiler can optimize this if-else at compile time + const auto num_plans = kWrite ? num_write : num_compress; + const auto plan_ptr = kWrite ? write_plan : compress_plan; + if (global_pid >= num_plans) return; + + const auto& [ragged_id, global_bid, position, window_len] = plan_ptr[global_pid]; + const int64_t split_offset = global_sid * kTileDim; + + // kv score + const auto kv_score_buffer = static_cast(_kv_score_buffer); + + // kv input + const auto kv_score_input = static_cast(_kv_score_input); + const auto kv_src = kv_score_input + ragged_id * kElementSize + split_offset; + + // kv output + const auto kv_compressed_output = static_cast(_kv_compressed_output); + const auto kv_out = kv_compressed_output + ragged_id * kHeadDim + split_offset; + + if (ragged_id == 0xFFFFFFFF) [[unlikely]] + return; + + // score bias (ape) + const auto score_bias = static_cast(_score_bias) + split_offset; + const auto seq_len = position + 1; + const int32_t index = indices[global_bid]; + + PDLWaitPrimary(); + + if constexpr (kMode == PageMode::Page4Align) { + const auto write_second_page = index; + const auto [load_first_page, load_second_page, write_first_page, last_pos] = extra[global_bid]; + if constexpr (kWrite) { + int32_t index; + if (position < static_cast(last_pos)) { + index = write_first_page; + } else { + index = write_second_page; + } + const auto kv_buf = kv_score_buffer + index * (kElementSize * 4) + split_offset; + c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 4); + } else { + int32_t index_overlap, index_normal; + if (window_len <= 4) { + index_overlap = load_second_page; + index_normal = load_second_page; // not used + } else { + index_overlap = load_first_page; + index_normal = load_second_page; + } + const auto kv_buf = kv_score_buffer + index_normal * (kElementSize * 4) + split_offset; + const auto kv_overlap = kv_score_buffer + index_overlap * (kElementSize * 4) + split_offset; + c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, window_len, kv_overlap); + } + } else { + static_assert(kMode == PageMode::RingBuffer, "Unsupported PageMode"); + const auto kv_buf = kv_score_buffer + index * (kElementSize * 8) + split_offset; + if constexpr (kWrite) { + c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 8); + } else { + c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, window_len); + } + } + + PDLTriggerSecondary(); +} + +template +struct FlashCompress4Kernel { + template + static constexpr auto decode_kernel = flash_c4_decode; + template + static constexpr auto prefill_kernel = flash_c4_prefill; + template + static constexpr auto prefill_c_kernel = prefill_kernel; + template + static constexpr auto prefill_w_kernel = prefill_kernel; + static constexpr uint32_t kBlockSize = 128; + static constexpr uint32_t kTileDim = kTileElements * device::kWarpThreads; + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static constexpr uint32_t kWarpsPerBlock = kBlockSize / device::kWarpThreads; + + using Self = FlashCompress4Kernel; + + static void run_decode( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::Optional extra) { + using namespace host; + + // this should not happen in practice + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + const auto extra_ptr = _get_extra_pointer(B, device_, extra); + const auto page_size = extra_ptr != nullptr ? 4 : 8; + + TensorMatcher({-1, page_size, kHeadDim * 4}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({B, kHeadDim * 4}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({B, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({8, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + TensorMatcher({B}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + TensorMatcher({B}) // seq lens + .with_dtype() + .with_device(device_) + .verify(seq_lens); + + const auto device = device_.unwrap(); + const auto batch_size = static_cast(B.unwrap()); + const auto params = Compress4DecodeParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .extra = static_cast(extra_ptr), + .batch_size = batch_size, + }; + const auto kernel = extra_ptr != nullptr ? decode_kernel // + : decode_kernel; + const uint32_t num_blocks = div_ceil(batch_size * kNumSplit, kWarpsPerBlock); + LaunchKernel(num_blocks, kBlockSize, device) // + .enable_pdl(kUsePDL)(kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_score_buffer, + const tvm::ffi::TensorView kv_score_input, + const tvm::ffi::TensorView kv_compressed_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView indices, + const tvm::ffi::TensorView compress_plan, + const tvm::ffi::TensorView write_plan, + const tvm::ffi::Optional extra) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto N = SymbolicSize{"num_q_tokens"}; + auto X = SymbolicSize{"compress_tokens"}; + auto Y = SymbolicSize{"write_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + const auto extra_ptr = _get_extra_pointer(B, device_, extra, /*is_prefill=*/true); + const auto page_size = extra_ptr != nullptr ? 4 : 8; + + TensorMatcher({-1, page_size, kHeadDim * 4}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_score_buffer); + TensorMatcher({N, kHeadDim * 4}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_score_input); + TensorMatcher({N, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_compressed_output); + TensorMatcher({8, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + TensorMatcher({B}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + TensorMatcher({X, compress::kPrefillPlanDim}) // compress plan + .with_dtype() + .with_device(device_) + .verify(compress_plan); + TensorMatcher({Y, compress::kPrefillPlanDim}) // write plan + .with_dtype() + .with_device(device_) + .verify(write_plan); + + const auto device = device_.unwrap(); + const auto batch_size = static_cast(B.unwrap()); + const auto num_q_tokens = static_cast(N.unwrap()); + const auto num_c = static_cast(X.unwrap()); + const auto num_w = static_cast(Y.unwrap()); + const auto params = Compress4PrefillParams{ + .kv_score_buffer = kv_score_buffer.data_ptr(), + .kv_score_input = kv_score_input.data_ptr(), + .kv_compressed_output = kv_compressed_output.data_ptr(), + .score_bias = ape.data_ptr(), + .indices = static_cast(indices.data_ptr()), + .extra = static_cast(extra_ptr), + .compress_plan = static_cast(compress_plan.data_ptr()), + .write_plan = static_cast(write_plan.data_ptr()), + .num_compress = num_c, + .num_write = num_w, + }; + RuntimeCheck(num_q_tokens >= batch_size, "num_q_tokens must be >= batch_size"); + RuntimeCheck(num_q_tokens >= std::max(num_c, num_w), "invalid prefill plan"); + if (const auto num_c_blocks = div_ceil(num_c * kNumSplit, kWarpsPerBlock)) { + const auto c_kernel = extra_ptr != nullptr ? prefill_c_kernel // + : prefill_c_kernel; + LaunchKernel(num_c_blocks, kBlockSize, device) // + .enable_pdl(kUsePDL)(c_kernel, params); + } + if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerBlock)) { + const auto w_kernel = extra_ptr != nullptr ? prefill_w_kernel // + : prefill_w_kernel; + LaunchKernel(num_w_blocks, kBlockSize, device) // + .enable_pdl(kUsePDL)(w_kernel, params); + } + } + + // some auxiliary functions + private: + static const void* _get_extra_pointer( + host::SymbolicSize& B, // batch_size + host::SymbolicDevice& device, + const tvm::ffi::Optional& extra, + bool is_prefill = false) { + // only have value when using page-aligned mode + if (!extra.has_value()) return nullptr; + const auto& extra_tensor = extra.value(); + /// NOTE: the metadata layout is different for prefill and decode: + /// for prefill, last 4 are: + /// load overlap | load normal | write overlap | last written page + /// for decode, last 1 is the write (also load) overlap + host::TensorMatcher({B, is_prefill ? 4 : 1}) // extra tensor + .with_dtype() + .with_device(device) + .verify(extra_tensor); + const auto data_ptr = extra_tensor.data_ptr(); + host::RuntimeCheck(data_ptr != nullptr, "extra tensor data ptr is null"); + if (is_prefill) { + static_assert(alignof(C4IndexBundle) == 16); + host::RuntimeCheck(std::bit_cast(data_ptr) % 16 == 0, "extra tensor is not properly aligned"); + } + return data_ptr; + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh new file mode 100644 index 0000000000..efa9f05100 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh @@ -0,0 +1,405 @@ +/** + * \brief Here's some dimension info for the main buffer used in C4 prefill and decode. + * + * kv_buffer: [num_indices, 8, head_dim * 4] + * - last dimension layout: | kv overlap | kv | score overlap | score | + * kv_input: [batch_size, head_dim * 4] + * kv_output: [batch_size, head_dim] + * score_bias (ape): [8, head_dim] + * plan_c/plan_w: [variable length] + * + * For prefill, batch_size = num_q_tokens + */ + +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include +#include + +namespace { + +using PlanD = device::compress::DecodePlan; +using PlanC = device::compress::CompressPlan; +using PlanW = device::compress::WritePlan; + +/// \brief Each thread will handle this many elements (split along head_dim) +constexpr int32_t kTileElements = 4; + +/// \brief Need to improve register usage to reduce latency +#define C4_KERNEL __global__ __launch_bounds__(128, 4) +#define WRITE_KERNEL __global__ __launch_bounds__(128, 16) + +struct Compress4DecodeParams { + void* __restrict__ kv_buffer; + const void* __restrict__ kv_input; + void* __restrict__ kv_output; + const void* __restrict__ score_bias; + const PlanD* __restrict__ plan_d; + uint32_t batch_size; +}; + +struct Compress4PrefillParams { + void* __restrict__ kv_buffer; + const void* __restrict__ kv_input; + void* __restrict__ kv_output; + const void* __restrict__ score_bias; + const PlanC* __restrict__ plan_c; + const PlanW* __restrict__ plan_w; + uint32_t num_compress; + uint32_t num_write; +}; + +template +struct C4Trait { + static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 128 + static constexpr int64_t kHeadDim = kHeadDim_; + static constexpr int64_t kOverlapOffset = kHeadDim; + static constexpr int64_t kScoreOffset = kHeadDim * 2; + static constexpr int64_t kElementSize = kHeadDim * 4; + static constexpr int64_t kPageElementSize = 4 * kElementSize; // page size = 4 + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static_assert(kHeadDim % kTileDim == 0); +}; + +template +SGL_DEVICE void c4_forward( + const InFloat* kv_buf_0, // overlap [4n - 4, 4n - 1] + const InFloat* kv_buf_1, // normal [4n + 0, 4n + 3] + const InFloat* kv_src, // ragged pointer at position = 4n + 3 + OutFloat* kv_out, + const InFloat* score_bias, + const bool should_overlap, + const int32_t buffer_len) { + using namespace device; + + /// NOTE: part 1: load kv + score + using StorageIn = AlignedVector; + /// NOTE: load one tile_dim (< head_dim) at at time + const auto gmem_in = tile::Memory::warp(); + StorageIn kv[8]; + StorageIn score[8]; + StorageIn bias[8]; + +#pragma unroll + for (int32_t i = 0; i < 8; ++i) { + bias[i] = gmem_in.load(score_bias + i * Trait::kHeadDim); + } + + if (should_overlap) { + const auto kv_start = kv_src - 7 * Trait::kElementSize; // point to start +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + const auto src = i < buffer_len ? kv_buf_0 : kv_start; + const auto base = src + i * Trait::kElementSize; + kv[i] = gmem_in.load(base); + score[i] = gmem_in.load(base + Trait::kScoreOffset); + } + } else { + [[unlikely]]; + constexpr float kFloatNegInf = -FLT_MAX; +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + kv[i].fill(cast(0.0f)); + score[i].fill(cast(kFloatNegInf)); + } + } + + const auto kv_start = kv_src - 3 * Trait::kElementSize; // point to start +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + const auto src = i + 4 < buffer_len ? kv_buf_1 : kv_start; + const auto base = src + i * Trait::kElementSize + Trait::kOverlapOffset; + kv[i + 4] = gmem_in.load(base); + score[i + 4] = gmem_in.load(base + Trait::kScoreOffset); + } + + /// NOTE: part 2: safe online softmax + weighted sum + using StorageOut = AlignedVector; + const auto gmem_out = tile::Memory::warp(); + StorageOut result; + + // consume 32 fp registers + float score_fp32[kTileElements][8]; + + // convert to fp32 and apply bias first +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + for (int32_t j = 0; j < 8; ++j) { + score_fp32[i][j] = cast(score[j][i]) + cast(bias[j][i]); + } + } + +#pragma unroll + for (int32_t i = 0; i < kTileElements; ++i) { + const auto& score = score_fp32[i]; + float max_value = score[0]; + float sum_exp_value = 0.0f; + +#pragma unroll + for (int32_t j = 1; j < 8; ++j) { + const auto fp32_score = score[j]; + max_value = fmaxf(max_value, fp32_score); + } + + float sum_product = 0.0f; +#pragma unroll + for (int32_t j = 0; j < 8; ++j) { + const auto fp32_score = score[j]; + const auto exp_score = expf(fp32_score - max_value); + sum_product += cast(kv[j][i]) * exp_score; + sum_exp_value += exp_score; + } + + result[i] = cast(sum_product / sum_exp_value); + } + + // overlap the store with the next iteration's load + PDLTriggerSecondary(); + gmem_out.store(kv_out, result); +} + +template +SGL_DEVICE void c4_write_decode(InFloat* kv_buf, const InFloat* kv_src) { + using namespace device; + + using StorageIn = AlignedVector; + const auto gmem = tile::Memory::warp(); + + StorageIn data[4]; +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + data[i] = gmem.load(kv_src + Trait::kHeadDim * i); + } +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + gmem.store(kv_buf + Trait::kHeadDim * i, data[i]); + } +} + +template +C4_KERNEL void flash_c4_decode(const __grid_constant__ Compress4DecodeParams params) { + using namespace device; + using Trait = C4Trait; + + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_bid = global_wid / Trait::kNumSplit; // batch id + const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id + const int64_t split_offset = global_sid * Trait::kTileDim; + if (global_bid >= params.batch_size) return; + + const auto plan = params.plan_d[global_bid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_output = static_cast(params.kv_output) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + const auto score_bias = static_cast(params.score_bias) + split_offset; + + const auto kv_src = kv_input + global_bid * Trait::kElementSize; + const auto kv_out = kv_output + global_bid * Trait::kHeadDim; + const auto kv_buf_0 = kv_buffer + plan.read_page_0 * Trait::kPageElementSize; + const auto kv_buf_1 = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; + const auto kv_dst = kv_buffer + plan.write_loc * Trait::kElementSize; + + PDLWaitPrimary(); + c4_write_decode(kv_dst, kv_src); + if (plan.seq_len % 4 == 0) { + const auto need_overlap = plan.seq_len > 4; + c4_forward(kv_buf_0, kv_buf_1, kv_src, kv_out, score_bias, need_overlap, 8); + } +} + +template +C4_KERNEL void flash_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { + using namespace device; + using Trait = C4Trait; + + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id + const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id + const int64_t split_offset = global_sid * Trait::kTileDim; + if (global_pid >= params.num_compress) return; + + const auto plan = params.plan_c[global_pid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_output = static_cast(params.kv_output) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + const auto score_bias = static_cast(params.score_bias) + split_offset; + if (plan.is_invalid()) return; + + const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; + // Compact output: one row per compress plan, indexed by `global_pid`. + const auto kv_out = kv_output + global_pid * Trait::kHeadDim; + const auto kv_buf_0 = kv_buffer + plan.read_page_0 * Trait::kPageElementSize; + const auto kv_buf_1 = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; + const bool need_overlap = plan.seq_len > 4; + PDLWaitPrimary(); + c4_forward(kv_buf_0, kv_buf_1, kv_src, kv_out, score_bias, need_overlap, plan.buffer_len); +} + +template +WRITE_KERNEL void write_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { + using namespace device; + using Trait = C4Trait; + using StorageIn = AlignedVector; + + const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t global_wid = global_tid / kWarpThreads; // warp id + const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id + const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id + // split the contiguous `kHeadDim * 4` into `kNumSplit` tiles + // each warp handles 1 contiguous tile (in contrast, decode handle the strided head_dim) + const int64_t split_offset = global_sid * (Trait::kTileDim * 4); + if (global_pid >= params.num_write) return; + + const auto plan = params.plan_w[global_pid]; + const auto kv_input = static_cast(params.kv_input) + split_offset; + const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; + if (plan.is_invalid()) return; + + // each warp will handle a contiguous region + const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; + const auto kv_buf = kv_buffer + plan.write_loc * Trait::kElementSize; + const auto gmem = tile::Memory::warp(); + + PDLWaitPrimary(); + StorageIn data[4]; +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + data[i] = gmem.load(kv_src, i); + } + PDLTriggerSecondary(); +#pragma unroll + for (int32_t i = 0; i < 4; ++i) { + gmem.store(kv_buf, data[i], i); + } +} + +template +struct FlashCompress4Kernel { + static constexpr auto decode_kernel = flash_c4_decode; + static constexpr auto prefill_c_kernel = flash_c4_prefill; + static constexpr auto prefill_w_kernel = write_c4_prefill; + static constexpr uint32_t kBlockSize = 128; + static constexpr uint32_t kTileDim = kTileElements * device::kWarpThreads; + static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; + static constexpr uint32_t kWarpsPerBlock = kBlockSize / device::kWarpThreads; + using Trait = C4Trait; + + static void run_decode( + const tvm::ffi::TensorView kv_buffer, + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_d_) { + using namespace host; + + auto N = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 4, Trait::kElementSize}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_buffer); + TensorMatcher({N, Trait::kElementSize}) // kv score input + .with_dtype() + .with_device(device_) + .verify(kv_input); + TensorMatcher({N, kHeadDim}) // kv compressed output + .with_dtype() + .with_device(device_) + .verify(kv_output); + TensorMatcher({8, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + + const auto plan_d = compress::verify_plan_d(plan_d_, N, device_); + const auto batch_size = static_cast(N.unwrap()); + const auto params = Compress4DecodeParams{ + .kv_buffer = kv_buffer.data_ptr(), + .kv_input = kv_input.data_ptr(), + .kv_output = kv_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_d = plan_d, + .batch_size = batch_size, + }; + const uint32_t num_blocks = div_ceil(batch_size * kNumSplit, kWarpsPerBlock); + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(decode_kernel, params); + } + + static void run_prefill( + const tvm::ffi::TensorView kv_buffer, + const tvm::ffi::TensorView kv_input, + const tvm::ffi::TensorView kv_output, + const tvm::ffi::TensorView ape, + const tvm::ffi::TensorView plan_c_, + const tvm::ffi::TensorView plan_w_) { + using namespace host; + + auto N = SymbolicSize{"num_q_tokens"}; + auto C = SymbolicSize{"num_c_plans"}; + auto W = SymbolicSize{"num_w_plans"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({-1, 4, Trait::kElementSize}) // kv score + .with_dtype() + .with_device(device_) + .verify(kv_buffer); + TensorMatcher({N, Trait::kElementSize}) // kv score input (ragged) + .with_dtype() + .with_device(device_) + .verify(kv_input); + TensorMatcher({C, kHeadDim}) // kv compressed output (compact) + .with_dtype() + .with_device(device_) + .verify(kv_output); + TensorMatcher({8, kHeadDim}) // ape + .with_dtype() + .with_device(device_) + .verify(ape); + const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); + const auto plan_w = compress::verify_plan_w(plan_w_, W, device_); + const auto device = device_.unwrap(); + const auto num_q_tokens = static_cast(N.unwrap()); + const auto num_c = static_cast(C.unwrap()); + const auto num_w = static_cast(W.unwrap()); + const auto params = Compress4PrefillParams{ + .kv_buffer = kv_buffer.data_ptr(), + .kv_input = kv_input.data_ptr(), + .kv_output = kv_output.data_ptr(), + .score_bias = ape.data_ptr(), + .plan_c = plan_c, + .plan_w = plan_w, + .num_compress = num_c, + .num_write = num_w, + }; + RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); + if (const auto num_c_blocks = div_ceil(num_c * kNumSplit, kWarpsPerBlock)) { + LaunchKernel(num_c_blocks, kBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_c_kernel, params); + } + if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerBlock)) { + LaunchKernel(num_w_blocks, kBlockSize, device) // + .enable_pdl(kUsePDL)(prefill_w_kernel, params); + } + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh new file mode 100644 index 0000000000..3e4aaaf5f0 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh @@ -0,0 +1,839 @@ +#include +#include +#include + +#include +#include + +#include + +#include +#include + +#include +#include + +namespace host::compress { + +constexpr auto kDLUInt8 = DLDataType{.code = kDLUInt, .bits = 8, .lanes = 1}; + +using PlanC = CompressPlan; +using PlanW = WritePlan; +using PlanD = DecodePlan; + +using RID_T = int64_t; +using R2T_T = int32_t; +using F2S_T = int64_t; +using IDX_T = int64_t; + +/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 +SGL_DEVICE __host__ PlanW pack_w(uint32_t ragged_id, uint32_t batch_id, int32_t seq_len) { + return {static_cast(ragged_id | batch_id << 16), seq_len}; +} + +/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 +SGL_DEVICE uint2 unpack_w(PlanW plan) { + return {static_cast(plan.ragged_id), static_cast(plan.ragged_id >> 16)}; +} + +struct Prefill0Params { + PlanC* plan_c; + PlanW* plan_w; + const IDX_T* seq_lens_ptr; // [batch_size] + const IDX_T* extend_lens_ptr; // [batch_size] + uint32_t batch_size; + uint32_t num_q_tokens; + int32_t compress_ratio; + int32_t swa_page_size; + int32_t mtp_pad; +}; + +struct Prefill1Params { + PlanC* plan_c; + PlanW* plan_w; + const RID_T* rid_ptr; // [batch_size] + const R2T_T* r2t_ptr; // [num_reqs, stride_r2t] + const F2S_T* f2s_ptr; // [num_swa_slots] + int64_t stride_r2t; + uint32_t num_c; + uint32_t num_w; + uint32_t num_c_padded; + uint32_t num_w_padded; + uint32_t num_work; + int32_t swa_page_size; + int32_t ring_size; + int32_t compress_ratio; +}; + +struct DecodeParams { + PlanD* plan_d; + const RID_T* rid_ptr; // [batch_size] + const R2T_T* r2t_ptr; // [num_reqs, stride_r2t] + const F2S_T* f2s_ptr; // [num_swa_slots] + const IDX_T* seq_ptr; // [batch_size] + int64_t stride_r2t; + uint32_t batch_size; + int32_t swa_page_size; + int32_t ring_size; + int32_t compress_ratio; +}; + +struct Prefill1ParamsLegacy { + PlanC* plan_c; + PlanW* plan_w; + const RID_T* rid_ptr; // [batch_size] + uint32_t num_c; + uint32_t num_w; + uint32_t num_c_padded; + uint32_t num_w_padded; + uint32_t num_work; + int32_t compress_ratio; +}; + +struct DecodeParamsLegacy { + PlanD* plan_d; + const RID_T* rid_ptr; // [batch_size] + const IDX_T* seq_ptr; // [batch_size] + uint32_t batch_size; + int32_t compress_ratio; +}; + +inline constexpr uint32_t kMaxPrefillBatchSize = 1024; + +SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { + static_assert(device::kWarpThreads == 32); +#pragma unroll + for (uint32_t offset = 1; offset < 32; offset *= 2) { +#ifndef USE_ROCM + uint32_t n = __shfl_up_sync(device::kFullMask, val, offset); +#else + uint32_t n = __shfl_up(val, offset, 32); +#endif + if (lane_id >= offset) val += n; + } + return val; +} + +/// Warp-wide max/min for integer types. `device::warp::reduce_max` routes through +/// `dtype_trait::max` which is only specialized for FP types. +SGL_DEVICE uint32_t warp_reduce_max_u32(uint32_t val) { +#pragma unroll + for (uint32_t mask = 16; mask > 0; mask >>= 1) { +#ifndef USE_ROCM + val = max(val, __shfl_xor_sync(device::kFullMask, val, mask, 32)); +#else + val = max(val, __shfl_xor(val, mask, 32)); +#endif + } + return val; +} + +SGL_DEVICE uint32_t warp_reduce_min_u32(uint32_t val) { +#pragma unroll + for (uint32_t mask = 16; mask > 0; mask >>= 1) { +#ifndef USE_ROCM + val = min(val, __shfl_xor_sync(device::kFullMask, val, mask, 32)); +#else + val = min(val, __shfl_xor(val, mask, 32)); +#endif + } + return val; +} + +__global__ __launch_bounds__(1024, 1) // + void plan_compress_prefill_kernel0(const Prefill0Params params) { + using namespace device; + const auto tx = threadIdx.x; + const auto block_size = kMaxPrefillBatchSize; + constexpr auto kNumWarps = kMaxPrefillBatchSize / kWarpThreads; + const auto cr = params.compress_ratio; + const auto sps = params.swa_page_size; + const bool is_overlap = (cr == 4); + const int32_t window_size = cr * (is_overlap ? 2 : 1); + + alignas(128) __shared__ uint32_t counter_c; + alignas(128) __shared__ uint32_t counter_w; + __shared__ int32_t s_seq_len[kMaxPrefillBatchSize]; + __shared__ int32_t s_prefix_len[kMaxPrefillBatchSize]; + __shared__ uint32_t warp_max[kNumWarps]; + __shared__ uint32_t warp_min[kNumWarps]; + __shared__ uint32_t s_max_extend; + __shared__ uint32_t s_min_extend; + + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + // === Stage A: load per-batch fields, init shared scratch === + int32_t seq_len = 0, extend_len = 0, prefix_len = 0; + if (tx < params.batch_size) { + seq_len = static_cast(params.seq_lens_ptr[tx]); + extend_len = static_cast(params.extend_lens_ptr[tx]); + prefix_len = seq_len - extend_len; + s_seq_len[tx] = seq_len; + s_prefix_len[tx] = prefix_len; + } + if (tx == 0) { + counter_c = 0; + counter_w = 0; + } + if (tx < kNumWarps) { + warp_max[tx] = 0; + warp_min[tx] = 0xFFFFFFFFu; + } + + // === Stage B: min/max(extend_len) for MTP-uniform detection === + // For min, treat threads outside `batch_size` as +inf so they don't pull the min down. + const uint32_t e_for_max = static_cast(extend_len); + const uint32_t e_for_min = (tx < params.batch_size) ? e_for_max : 0xFFFFFFFFu; + warp_max[warp_id] = warp_reduce_max_u32(e_for_max); + warp_min[warp_id] = warp_reduce_min_u32(e_for_min); + __syncthreads(); + if (warp_id == 0) { + s_max_extend = warp_reduce_max_u32(warp_max[lane_id]); + s_min_extend = warp_reduce_min_u32(warp_min[lane_id]); + } + __syncthreads(); + + const auto num_q = params.num_q_tokens; + // MTP-uniform: every batch shares the same small extend_len `E`, so we can decompose + // a global token id `k` into (batch_id, j) = (k / E, k % E) and skip the per-batch loop. + const bool is_mtp_extend = (s_min_extend == s_max_extend) && (s_max_extend > 0) && (s_max_extend <= 32); + + // === Stage C: emit valid plans, slot allocation via shared-mem atomicAdd === + if (is_mtp_extend) { + // Path 1: token-driven. Each global token id maps to exactly one (batch_id, j). + const uint32_t E = s_max_extend; + for (uint32_t k = tx; k < num_q; k += block_size) { + const uint32_t batch_id = k / E; + const uint32_t j = k % E; + const int32_t pl = s_prefix_len[batch_id]; + const int32_t sl = s_seq_len[batch_id]; + const int32_t position = pl + static_cast(j); + const uint32_t ragged_id = k; + + if ((position + 1) % cr == 0) { + const int32_t buffer_len = window_size - min(static_cast(j) + 1, window_size); + const uint32_t out_idx = atomicAdd(&counter_c, 1u); + params.plan_c[out_idx] = { + .seq_len = static_cast(position + 1), + .ragged_id = static_cast(ragged_id), + .buffer_len = static_cast(buffer_len), + .read_page_0 = -1, + .read_page_1 = static_cast(batch_id), + }; + } + + const int32_t last_c_pos = (sl / cr) * cr; + const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); + bool do_write = position >= first_w_pos; + if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); + if (do_write) { + const uint32_t out_idx = atomicAdd(&counter_w, 1u); + params.plan_w[out_idx] = pack_w(ragged_id, batch_id, position + 1); + } + } + } else { + // Path 2: general prefill (long extend_len). Iterate batches in an outer loop; + // the whole block sweeps each batch's tokens in parallel. + uint32_t base_e = 0; + for (uint32_t batch_id = 0; batch_id < params.batch_size; ++batch_id) { + const int32_t pl = s_prefix_len[batch_id]; + const int32_t sl = s_seq_len[batch_id]; + const int32_t el = sl - pl; + const int32_t last_c_pos = (sl / cr) * cr; + const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); + for (int32_t j = static_cast(tx); j < el; j += static_cast(block_size)) { + const int32_t position = pl + j; + const uint32_t ragged_id = base_e + static_cast(j); + + if ((position + 1) % cr == 0) { + const int32_t buffer_len = window_size - min(j + 1, window_size); + const uint32_t out_idx = atomicAdd(&counter_c, 1u); + params.plan_c[out_idx] = { + .seq_len = static_cast(position + 1), + .ragged_id = static_cast(ragged_id), + .buffer_len = static_cast(buffer_len), + .read_page_0 = -1, + .read_page_1 = static_cast(batch_id), + }; + } + + bool do_write = position >= first_w_pos; + if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); + if (do_write) { + const uint32_t out_idx = atomicAdd(&counter_w, 1u); + params.plan_w[out_idx] = pack_w(ragged_id, static_cast(batch_id), position + 1); + } + } + base_e += static_cast(el); + } + } + __syncthreads(); + + // === Stage D: pad [counter_c, num_q) / [counter_w, num_q) with invalid === + const auto total_c = counter_c; + const auto total_w = counter_w; + for (uint32_t k = total_c + tx; k < num_q; k += block_size) { + params.plan_c[k] = PlanC::invalid(); + } + for (uint32_t k = total_w + tx; k < num_q; k += block_size) { + params.plan_w[k] = PlanW::invalid(); + } +} + +/// NOTE: stage 1 +__global__ void plan_compress_prefill_kernel_1(const Prefill1Params params) { + const auto idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= params.num_work) return; + auto plan_c = idx < params.num_c ? params.plan_c[idx] : PlanC::invalid(); + auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid(); + + const auto compute_loc = [&](int32_t swa_loc) { + const auto swa_page = swa_loc / params.swa_page_size; + const auto ring_offset = swa_loc % params.ring_size; + return swa_page * params.ring_size + ring_offset; + }; + + if (!plan_c.is_invalid()) { // 1. in bound. 2. not masked + if (plan_c.buffer_len > 0) { + const auto batch_id = plan_c.read_page_1; + const auto rid = params.rid_ptr[batch_id]; + const auto mapping = params.r2t_ptr + rid * params.stride_r2t; + // `seq_len` should be ratio-aligned here + const auto position_1 = static_cast(plan_c.seq_len - 1); + // only used for c4, harmless for c128 + const auto position_0 = max(position_1 - params.compress_ratio, 0); + const auto raw_loc_0 = mapping[position_0]; + const auto raw_loc_1 = mapping[position_1]; + const auto swa_loc_0 = params.f2s_ptr[raw_loc_0]; + const auto swa_loc_1 = params.f2s_ptr[raw_loc_1]; + plan_c.read_page_0 = compute_loc(swa_loc_0) / params.compress_ratio; + plan_c.read_page_1 = compute_loc(swa_loc_1) / params.compress_ratio; + params.plan_c[idx] = plan_c; + } + } else if (idx < params.num_c_padded) { + params.plan_c[idx] = PlanC::invalid(); + } + + if (!plan_w.is_invalid()) { // 1. in bound. 2. not masked + const auto [ragged_id, batch_id] = unpack_w(plan_w); + const auto rid = params.rid_ptr[batch_id]; + const auto mapping = params.r2t_ptr + rid * params.stride_r2t; + // `seq_len` (`write_loc`) may not be aligned here + const auto position = static_cast(plan_w.write_loc - 1); + const auto raw_loc = mapping[position]; + const auto swa_loc = params.f2s_ptr[raw_loc]; + plan_w.ragged_id = ragged_id; + plan_w.write_loc = compute_loc(swa_loc); + params.plan_w[idx] = plan_w; + } else if (idx < params.num_w_padded) { + params.plan_w[idx] = PlanW::invalid(); + } +} + +__global__ void plan_compress_decode_kernel(const DecodeParams params) { + const auto idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= params.batch_size) return; + const auto rid = params.rid_ptr[idx]; + const auto mapping = params.r2t_ptr + rid * params.stride_r2t; + const auto compute_loc = [&](int32_t swa_loc) { + const auto swa_page = swa_loc / params.swa_page_size; + const auto ring_offset = swa_loc % params.ring_size; + return swa_page * params.ring_size + ring_offset; + }; + const auto seq_len = static_cast(params.seq_ptr[idx]); + const auto position_1 = static_cast(seq_len - 1); + const auto position_0 = max(position_1 - params.compress_ratio, 0); + const auto raw_loc_0 = mapping[position_0]; + const auto raw_loc_1 = mapping[position_1]; + const auto swa_loc_0 = params.f2s_ptr[raw_loc_0]; + const auto swa_loc_1 = params.f2s_ptr[raw_loc_1]; + const auto write_loc = compute_loc(swa_loc_1); + const auto read_page_0 = compute_loc(swa_loc_0) / params.compress_ratio; + const auto read_page_1 = write_loc / params.compress_ratio; + params.plan_d[idx] = { + .seq_len = static_cast(seq_len), + .write_loc = write_loc, + .read_page_0 = read_page_0, + .read_page_1 = read_page_1, + }; +} + +__global__ void plan_compress_prefill_legacy_kernel(const Prefill1ParamsLegacy params) { + const auto idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= params.num_work) return; + auto plan_c = idx < params.num_c ? params.plan_c[idx] : PlanC::invalid(); + auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid(); + + /// Per-request ring buffer slot translation: + /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 + /// - c128: page = rid; slot = rid * 128 + position % 128 + const auto legacy_compute_page = [&](int32_t rid, int32_t position) { + if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); + return rid; // c128 + }; + const auto legacy_compute_loc = [&](int32_t rid, int32_t position) { + const auto remainder = position % params.compress_ratio; + return legacy_compute_page(rid, position) * params.compress_ratio + remainder; + }; + + if (!plan_c.is_invalid()) { + const auto batch_id = plan_c.read_page_1; + const auto rid = static_cast(params.rid_ptr[batch_id]); + // `seq_len` is ratio-aligned for compress events + const auto position_1 = static_cast(plan_c.seq_len) - 1; + const auto position_0 = max(position_1 - params.compress_ratio, 0); + plan_c.read_page_0 = legacy_compute_page(rid, position_0); + plan_c.read_page_1 = legacy_compute_page(rid, position_1); + params.plan_c[idx] = plan_c; + } else if (idx < params.num_c_padded) { + params.plan_c[idx] = PlanC::invalid(); + } + + if (!plan_w.is_invalid()) { + const auto [ragged_id, batch_id] = unpack_w(plan_w); + const auto rid = static_cast(params.rid_ptr[batch_id]); + // `write_loc` carries (position + 1) at this stage; may not be ratio-aligned + const auto position = static_cast(plan_w.write_loc) - 1; + plan_w.ragged_id = ragged_id; + plan_w.write_loc = legacy_compute_loc(rid, position); + params.plan_w[idx] = plan_w; + } else if (idx < params.num_w_padded) { + params.plan_w[idx] = PlanW::invalid(); + } +} + +__global__ void plan_compress_decode_legacy_kernel(const DecodeParamsLegacy params) { + const auto idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= params.batch_size) return; + /// Per-request ring buffer slot translation: + /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 + /// - c128: page = rid; slot = rid * 128 + position % 128 + const auto legacy_compute_page = [&](int32_t rid, int32_t position) { + if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); + return rid; // c128 + }; + const auto legacy_compute_loc = [&](int32_t rid, int32_t position) { + const auto remainder = position % params.compress_ratio; + return legacy_compute_page(rid, position) * params.compress_ratio + remainder; + }; + const auto rid = static_cast(params.rid_ptr[idx]); + const auto seq_len = static_cast(params.seq_ptr[idx]); + const auto position_1 = seq_len - 1; + const auto position_0 = max(position_1 - params.compress_ratio, 0); + const auto write_loc = legacy_compute_loc(rid, position_1); + const auto read_page_0 = legacy_compute_page(rid, position_0); + const auto read_page_1 = legacy_compute_page(rid, position_1); + params.plan_d[idx] = { + .seq_len = static_cast(seq_len), + .write_loc = write_loc, + .read_page_0 = read_page_0, + .read_page_1 = read_page_1, + }; +} + +using PrefillPlan = tvm::ffi::Tuple; + +/** + * \brief Build c4/c128 prefill plan tensors. CPU-resident. + * Inputs (all CPU-resident): + * @param req_pool_indices `[batch_size]` int64_t + * @param req_to_token `[num_reqs, max_tokens_per_req]` int64_t + * @param full_to_swa `[num_swa_slots]` int64_t + * @param seq_lens `[batch_size]` int64 + * @param extend_lens `[batch_size]` int64 + * @param compress_plan `[num_q_tokens, 16]` uint8 (output) + * @param write_plan `[num_q_tokens, 8]` uint8 (output) + * @param compress_ratio 4 for c4, 128 for c128 + * @param use_cuda_graph Whether the plans will be used with cuda graph (affects padding) + * @return (compress plan tensor, write plan tensor) + */ +inline PrefillPlan plan_compress_prefill( + const tvm::ffi::TensorView req_pool_indices, // GPU + const tvm::ffi::TensorView req_to_token, // GPU + const tvm::ffi::TensorView full_to_swa, // GPU + const tvm::ffi::TensorView seq_lens, // CPU/GPU + const tvm::ffi::TensorView extend_lens, // CPU/GPU + const tvm::ffi::TensorView pin_buffer, // CPU + const uint32_t num_q_tokens, + const int32_t compress_ratio, + const int32_t swa_page_size, + const int32_t ring_size, + const bool use_cuda_graph) { + auto B = SymbolicSize{"batch_size"}; + auto N = SymbolicSize{"num_q_tokens"}; + auto cpu_or_gpu = SymbolicDevice{}; + auto device_ = SymbolicDevice{}; + cpu_or_gpu.set_options(); + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({-1, -1}) // + .with_dtype() + .with_device(device_) + .verify(req_to_token); + TensorMatcher({-1}) // + .with_dtype() + .with_device(device_) + .verify(full_to_swa); + TensorMatcher({B}) // + .with_dtype() + .with_device(cpu_or_gpu) + .verify(seq_lens) + .verify(extend_lens); + TensorMatcher({-1}) // + .with_dtype() + .with_device() + .verify(pin_buffer); + + const bool is_overlap = (compress_ratio == 4); + const int32_t window_size = compress_ratio * (is_overlap ? 2 : 1); + + const auto seq_ptr = static_cast(seq_lens.data_ptr()); + const auto ext_ptr = static_cast(extend_lens.data_ptr()); + const auto rid_ptr = static_cast(req_pool_indices.data_ptr()); + const auto r2t_ptr = static_cast(req_to_token.data_ptr()); + const auto f2s_ptr = static_cast(full_to_swa.data_ptr()); + + const auto batch_size = static_cast(B.unwrap()); + constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); + RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); + RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); + // `swa_page_size` >= `ring_size` >= `compress_ratio` + RuntimeCheck(swa_page_size % ring_size == 0 && ring_size % compress_ratio == 0); + + const auto device = device_.unwrap(); + const auto stream = LaunchKernel::resolve_device(device); + + constexpr int32_t kMaxMTPDraftTokens = 4; + const auto mtp_pad = std::min(ring_size - compress_ratio, kMaxMTPDraftTokens); + + if (cpu_or_gpu.unwrap().device_type == kDLGPU) { + // GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata directly + // on device, padding to num_q_tokens with invalid; kernel_1 then finalizes the + // SWA-translated read/write locations. Used for MTP / cuda-graph capture where + // a host sync would be expensive. + RuntimeCheck(batch_size <= kMaxPrefillBatchSize, "GPU plan only support batch size up to ", kMaxPrefillBatchSize); + auto C = ffi::empty({num_q_tokens, sizeof(PlanC)}, kDLUInt8, device); + auto W = ffi::empty({num_q_tokens, sizeof(PlanW)}, kDLUInt8, device); + const auto params0 = Prefill0Params{ + .plan_c = static_cast(C.data_ptr()), + .plan_w = static_cast(W.data_ptr()), + .seq_lens_ptr = seq_ptr, + .extend_lens_ptr = ext_ptr, + .batch_size = batch_size, + .num_q_tokens = num_q_tokens, + .compress_ratio = compress_ratio, + .swa_page_size = swa_page_size, + .mtp_pad = mtp_pad, + }; + LaunchKernel(1, kMaxPrefillBatchSize, device)(plan_compress_prefill_kernel0, params0); + // kernel_1 sees the already-padded buffers, so num_c == num_w == num_padded == num_q_tokens. + const auto params1 = Prefill1Params{ + .plan_c = static_cast(C.data_ptr()), + .plan_w = static_cast(W.data_ptr()), + .rid_ptr = rid_ptr, + .r2t_ptr = r2t_ptr, + .f2s_ptr = f2s_ptr, + .stride_r2t = req_to_token.stride(0), + .num_c = num_q_tokens, + .num_w = num_q_tokens, + .num_c_padded = num_q_tokens, + .num_w_padded = num_q_tokens, + .num_work = num_q_tokens, + .swa_page_size = swa_page_size, + .ring_size = ring_size, + .compress_ratio = compress_ratio, + }; + const auto block_size_1 = 256; + const auto num_blocks_1 = div_ceil(params1.num_work, block_size_1); + LaunchKernel(num_blocks_1, block_size_1, device)(plan_compress_prefill_kernel_1, params1); + return PrefillPlan{std::move(C), std::move(W)}; + } + + // CPU input path: only here do we need the pinned scratch buffer. + const auto pin_buffer_bytes = static_cast(pin_buffer.numel()) * sizeof(uint8_t); + RuntimeCheck(pin_buffer_bytes >= num_q_tokens * (sizeof(PlanC) + sizeof(PlanW))); + const auto plan_c_ptr = reinterpret_cast(pin_buffer.data_ptr()); + const auto plan_w_ptr = reinterpret_cast(plan_c_ptr + num_q_tokens); + + uint32_t counter = 0; + uint32_t counter_c = 0; + uint32_t counter_w = 0; + + const auto should_compress = [=](int32_t position) { return (position + 1) % compress_ratio == 0; }; + for (const auto i : irange(batch_size)) { + const int32_t seq_len = seq_ptr[i]; + const int32_t extend_len = ext_ptr[i]; + const int32_t prefix_len = seq_len - extend_len; + const int32_t last_c_pos = seq_len / compress_ratio * compress_ratio; + const int32_t first_w_pos = last_c_pos - (is_overlap ? compress_ratio : 0); + RuntimeCheck(0 < extend_len && extend_len <= seq_len); + const auto should_write = [=](int32_t position) { + if (position >= first_w_pos) return true; + return is_overlap && position % swa_page_size >= (swa_page_size - compress_ratio); + }; + for (const auto j : irange(extend_len)) { + const int32_t position = prefix_len + j; + const int32_t ragged_id = counter + j; + if (should_compress(position)) { + const auto buffer_len = window_size - std::min(j + 1, window_size); + plan_c_ptr[counter_c++] = { + .seq_len = static_cast(position + 1), + .ragged_id = static_cast(ragged_id), + .buffer_len = static_cast(buffer_len), + // to be filled by kernel + .read_page_0 = -1, + .read_page_1 = static_cast(i), + }; + } + if (should_write(position)) { + plan_w_ptr[counter_w++] = pack_w(ragged_id, i, position + 1); + } + } + counter += extend_len; + } + RuntimeCheck(counter == num_q_tokens); + + const auto copy_to_device = [stream](void* cuda_ptr, auto* host_ptr, size_t count) { + const auto size_bytes = count * sizeof(*host_ptr); + RuntimeDeviceCheck(cudaMemcpyAsync(cuda_ptr, host_ptr, size_bytes, cudaMemcpyHostToDevice, stream)); + }; + const auto num_c_padded = use_cuda_graph ? num_q_tokens : counter_c; + const auto num_w_padded = use_cuda_graph ? num_q_tokens : counter_w; + auto C = ffi::empty({num_c_padded, sizeof(PlanC)}, kDLUInt8, device); + auto W = ffi::empty({num_w_padded, sizeof(PlanW)}, kDLUInt8, device); + copy_to_device(C.data_ptr(), plan_c_ptr, counter_c); + copy_to_device(W.data_ptr(), plan_w_ptr, counter_w); + const auto params = Prefill1Params{ + .plan_c = static_cast(C.data_ptr()), + .plan_w = static_cast(W.data_ptr()), + .rid_ptr = rid_ptr, + .r2t_ptr = r2t_ptr, + .f2s_ptr = f2s_ptr, + .stride_r2t = req_to_token.size(1), + .num_c = counter_c, + .num_w = counter_w, + .num_c_padded = num_c_padded, + .num_w_padded = num_w_padded, + .num_work = std::max(num_c_padded, num_w_padded), + .swa_page_size = swa_page_size, + .ring_size = ring_size, + .compress_ratio = compress_ratio, + }; + const auto block_size = 256; + const auto num_blocks = div_ceil(params.num_work, block_size); + LaunchKernel(num_blocks, block_size, device)(plan_compress_prefill_kernel_1, params); + return PrefillPlan{std::move(C), std::move(W)}; +} + +inline tvm::ffi::Tensor plan_compress_decode( + const tvm::ffi::TensorView req_pool_indices, // GPU + const tvm::ffi::TensorView req_to_token, // GPU + const tvm::ffi::TensorView full_to_swa, // GPU + const tvm::ffi::TensorView seq_lens, // CPU/GPU + const int32_t compress_ratio, + const int32_t swa_page_size, + const int32_t ring_size) { + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({-1, -1}) // + .with_dtype() + .with_device(device_) + .verify(req_to_token); + TensorMatcher({-1}) // + .with_dtype() + .with_device(device_) + .verify(full_to_swa); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(seq_lens); + + const auto batch_size = static_cast(B.unwrap()); + const auto device = device_.unwrap(); + auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device); + const auto params = DecodeParams{ + .plan_d = static_cast(D.data_ptr()), + .rid_ptr = static_cast(req_pool_indices.data_ptr()), + .r2t_ptr = static_cast(req_to_token.data_ptr()), + .f2s_ptr = static_cast(full_to_swa.data_ptr()), + .seq_ptr = static_cast(seq_lens.data_ptr()), + .stride_r2t = req_to_token.size(1), + .batch_size = batch_size, + .swa_page_size = swa_page_size, + .ring_size = ring_size, + .compress_ratio = compress_ratio, + }; + const auto block_size = 256; + const auto num_blocks = div_ceil(batch_size, block_size); + LaunchKernel(num_blocks, block_size, device)(plan_compress_decode_kernel, params); + return D; +} + +/** + * \brief Build c4/c128 prefill plan tensors for the legacy non-paged ring + * buffer. Uses only `req_pool_indices` to derive ring slots: + * - c4 (overlap): each request occupies 2 contiguous pages (8 token slots) + * - c128: each request occupies 1 page (128 token slots) + * + * Inputs: + * @param req_pool_indices `[batch_size]` int64 (GPU) + * @param seq_lens `[batch_size]` int64 (CPU) + * @param extend_lens `[batch_size]` int64 (CPU) + * @param pin_buffer pinned scratch (CPU uint8) + * @return (compress plan tensor, write plan tensor) + */ +inline PrefillPlan plan_compress_prefill_legacy( + const tvm::ffi::TensorView req_pool_indices, // GPU + const tvm::ffi::TensorView seq_lens, // CPU + const tvm::ffi::TensorView extend_lens, // CPU + const tvm::ffi::TensorView pin_buffer, // CPU + const uint32_t num_q_tokens, + const int32_t compress_ratio, + const bool use_cuda_graph) { + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({B}) // + .with_dtype() + .with_device() + .verify(seq_lens) + .verify(extend_lens); + TensorMatcher({-1}) // + .with_dtype() + .with_device() + .verify(pin_buffer); + + const auto pin_buffer_bytes = static_cast(pin_buffer.numel()) * sizeof(uint8_t); + RuntimeCheck(pin_buffer_bytes >= num_q_tokens * (sizeof(PlanC) + sizeof(PlanW))); + const auto plan_c_ptr = reinterpret_cast(pin_buffer.data_ptr()); + const auto plan_w_ptr = reinterpret_cast(plan_c_ptr + num_q_tokens); + + const bool is_overlap = (compress_ratio == 4); + const auto seq_ptr = static_cast(seq_lens.data_ptr()); + const auto ext_ptr = static_cast(extend_lens.data_ptr()); + const auto rid_ptr = static_cast(req_pool_indices.data_ptr()); + + const auto window_size = compress_ratio * (is_overlap ? 2 : 1); + const auto batch_size = static_cast(B.unwrap()); + constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); + RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); + RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); + + uint32_t counter = 0; + uint32_t counter_c = 0; + uint32_t counter_w = 0; + const auto should_compress = [=](int32_t position) { return (position + 1) % compress_ratio == 0; }; + for (const auto i : irange(batch_size)) { + const int32_t seq_len = seq_ptr[i]; + const int32_t extend_len = ext_ptr[i]; + const int32_t prefix_len = seq_len - extend_len; + const int32_t last_c_pos = seq_len / compress_ratio * compress_ratio; + const int32_t first_w_pos = last_c_pos - (is_overlap ? compress_ratio : 0); + RuntimeCheck(0 < extend_len && extend_len <= seq_len); + const auto should_write = [=](int32_t position) { return position >= first_w_pos; }; + for (const auto j : irange(extend_len)) { + const int32_t position = prefix_len + j; + const int32_t ragged_id = counter + j; + if (should_compress(position)) { + const auto buffer_len = window_size - std::min(j + 1, window_size); + plan_c_ptr[counter_c++] = { + .seq_len = static_cast(position + 1), + .ragged_id = static_cast(ragged_id), + .buffer_len = static_cast(buffer_len), + // to be filled by kernel + .read_page_0 = -1, + .read_page_1 = static_cast(i), + }; + } + if (should_write(position)) { + plan_w_ptr[counter_w++] = pack_w(ragged_id, i, position + 1); + } + } + counter += extend_len; + } + RuntimeCheck(counter == num_q_tokens); + + const auto device = device_.unwrap(); + const auto stream = LaunchKernel::resolve_device(device); + const auto copy_to_device = [stream](void* cuda_ptr, auto* host_ptr, size_t count) { + const auto size_bytes = count * sizeof(*host_ptr); + RuntimeDeviceCheck(cudaMemcpyAsync(cuda_ptr, host_ptr, size_bytes, cudaMemcpyHostToDevice, stream)); + }; + const auto num_c_padded = use_cuda_graph ? num_q_tokens : counter_c; + const auto num_w_padded = use_cuda_graph ? num_q_tokens : counter_w; + auto C = ffi::empty({num_c_padded, sizeof(PlanC)}, kDLUInt8, device); + auto W = ffi::empty({num_w_padded, sizeof(PlanW)}, kDLUInt8, device); + copy_to_device(C.data_ptr(), plan_c_ptr, counter_c); + copy_to_device(W.data_ptr(), plan_w_ptr, counter_w); + const auto params = Prefill1ParamsLegacy{ + .plan_c = static_cast(C.data_ptr()), + .plan_w = static_cast(W.data_ptr()), + .rid_ptr = rid_ptr, + .num_c = counter_c, + .num_w = counter_w, + .num_c_padded = num_c_padded, + .num_w_padded = num_w_padded, + .num_work = std::max(num_c_padded, num_w_padded), + .compress_ratio = compress_ratio, + }; + const auto block_size = 256; + const auto num_blocks = div_ceil(params.num_work, block_size); + if (num_blocks > 0) { + LaunchKernel(num_blocks, block_size, device)(plan_compress_prefill_legacy_kernel, params); + } + return PrefillPlan{std::move(C), std::move(W)}; +} + +inline tvm::ffi::Tensor plan_compress_decode_legacy( + const tvm::ffi::TensorView req_pool_indices, // GPU + const tvm::ffi::TensorView seq_lens, // GPU + const int32_t compress_ratio) { + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(req_pool_indices); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(seq_lens); + RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); + + const auto batch_size = static_cast(B.unwrap()); + const auto device = device_.unwrap(); + auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device); + const auto params = DecodeParamsLegacy{ + .plan_d = static_cast(D.data_ptr()), + .rid_ptr = static_cast(req_pool_indices.data_ptr()), + .seq_ptr = static_cast(seq_lens.data_ptr()), + .batch_size = batch_size, + .compress_ratio = compress_ratio, + }; + const auto block_size = 256; + const auto num_blocks = div_ceil(batch_size, block_size); + LaunchKernel(num_blocks, block_size, device)(plan_compress_decode_legacy_kernel, params); + return D; +} + +} // namespace host::compress + +using namespace host::compress; // expose binding diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh new file mode 100644 index 0000000000..46acaa9c46 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh @@ -0,0 +1,208 @@ +#include +#include + +#include + +#include + +namespace host::compress { + +using PlanResult = tvm::ffi::Tuple; + +struct CompressParams { + PrefillPlan* __restrict__ compress_plan; + PrefillPlan* __restrict__ write_plan; + const int64_t* __restrict__ seq_lens; + const int64_t* __restrict__ extend_lens; + uint32_t batch_size; + uint32_t num_tokens; + uint32_t compress_ratio; + bool is_overlap; +}; + +inline constexpr uint32_t kBlockSize = 1024; + +#define PLAN_KERNEL __global__ __launch_bounds__(kBlockSize, 1) inline + +PLAN_KERNEL void plan_prefill_cuda(const __grid_constant__ CompressParams params) { + const auto &[ + compress_plan, write_plan, seq_lens, extend_lens, // pointers + batch_size, num_tokens, compress_ratio, is_overlap // values + ] = params; + + __shared__ uint32_t compress_counter; + __shared__ uint32_t write_counter; + + uint32_t batch_id = 0; + uint32_t counter = 0; + uint32_t extend_len = extend_lens[0]; + + const auto tid = threadIdx.x; + if (tid == 0) { + compress_counter = 0; + write_counter = 0; + } + __syncthreads(); + + for (uint32_t i = tid; i < num_tokens; i += blockDim.x) { + const uint32_t ragged_id = i; + uint32_t j = ragged_id - counter; + while (j >= extend_len) { + j -= extend_len; + batch_id += 1; + if (batch_id >= batch_size) [[unlikely]] + break; + counter += extend_len; + extend_len = extend_lens[batch_id]; + } + if (batch_id >= batch_size) [[unlikely]] + break; + const uint32_t seq_len = seq_lens[batch_id]; + const uint32_t extend_len = extend_lens[batch_id]; + const uint32_t prefix_len = seq_len - extend_len; + const uint32_t ratio = compress_ratio * (1 + is_overlap); + const uint32_t window_len = j + 1 < ratio ? ratio - (j + 1) : 0; + const uint32_t position = prefix_len + j; + const auto plan = PrefillPlan{ + .ragged_id = ragged_id, + .batch_id = batch_id, + .position = position, + .window_len = window_len, + }; + const uint32_t start_write_pos = [seq_len, compress_ratio, is_overlap] { + const uint32_t pos = seq_len / compress_ratio * compress_ratio; + if (!is_overlap) return pos; + return pos >= compress_ratio ? pos - compress_ratio : 0; + }(); + if ((position + 1) % compress_ratio == 0) { + const auto write_pos = atomicAdd(&compress_counter, 1); + compress_plan[write_pos] = plan; + } + if (position >= start_write_pos) { + const auto write_pos = atomicAdd(&write_counter, 1); + write_plan[write_pos] = plan; + } + } + __syncthreads(); + constexpr auto kInvalid = static_cast(-1); + const auto kInvalidPlan = PrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; + const auto compress_count = compress_counter; + const auto write_count = write_counter; + for (uint32_t i = compress_count + tid; i < num_tokens; i += blockDim.x) { + compress_plan[i] = kInvalidPlan; + } + for (uint32_t i = write_count + tid; i < num_tokens; i += blockDim.x) { + write_plan[i] = kInvalidPlan; + } +} + +inline PlanResult plan_prefill_host(const CompressParams& params, const bool use_cuda_graph) { + const auto &[ + compress_ptr, write_ptr, seq_lens_ptr, extend_lens_ptr, // pointers + batch_size, num_tokens, compress_ratio, is_overlap // values + ] = params; + + uint32_t counter = 0; + uint32_t compress_counter = 0; + uint32_t write_counter = 0; + const auto ratio = compress_ratio * (1 + is_overlap); + for (const auto i : irange(batch_size)) { + const uint32_t seq_len = seq_lens_ptr[i]; + const uint32_t extend_len = extend_lens_ptr[i]; + const uint32_t prefix_len = seq_len - extend_len; + RuntimeCheck(0 < extend_len && extend_len <= seq_len); + /// NOTE: `start_write_pos` must be a multiple of `compress_ratio` + const uint32_t start_write_pos = [seq_len, compress_ratio, is_overlap] { + const uint32_t pos = seq_len / compress_ratio * compress_ratio; + if (!is_overlap) return pos; + /// NOTE: to avoid unsigned integer underflow, don't use `pos - compress_ratio` + return pos >= compress_ratio ? pos - compress_ratio : 0; + }(); + /// NOTE: `position` is within [prefix_len, seq_len) + for (const auto j : irange(extend_len)) { + const uint32_t position = prefix_len + j; + const auto plan = PrefillPlan{ + .ragged_id = counter + j, + .batch_id = i, + .position = position, + .window_len = ratio - std::min(j + 1, ratio), + }; + RuntimeCheck(plan.is_valid(compress_ratio, is_overlap), "Internal error!"); + if ((position + 1) % compress_ratio == 0) { + compress_ptr[compress_counter++] = plan; + } + if (position >= start_write_pos) { + write_ptr[write_counter++] = plan; + } + } + counter += extend_len; + } + RuntimeCheck(counter == num_tokens, "input size ", counter, " != num_q_tokens ", num_tokens); + if (!use_cuda_graph) return PlanResult{compress_counter, write_counter}; + constexpr auto kInvalid = static_cast(-1); + constexpr auto kInvalidPlan = PrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; + for (const auto i : irange(compress_counter, num_tokens)) { + compress_ptr[i] = kInvalidPlan; + } + for (const auto i : irange(write_counter, num_tokens)) { + write_ptr[i] = kInvalidPlan; + } + return PlanResult{num_tokens, num_tokens}; +} + +inline PlanResult plan_prefill( + const tvm::ffi::TensorView extend_lens, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView compress_plan, + const tvm::ffi::TensorView write_plan, + const uint32_t compress_ratio, + const bool is_overlap, // for overlap transform, we have to keep 1 more extra window + const bool use_cuda_graph) { + auto N = SymbolicSize{"batch_size"}; + auto M = SymbolicSize{"num_tokens"}; + auto device = SymbolicDevice{}; + const bool is_cuda = [&] { + if (extend_lens.device().device_type == kDLCUDA) { + device.set_options(); + return true; + } else { + device.set_options(); + return false; + } + }(); + TensorMatcher({N}) // extend_lens and seq_lens + .with_dtype() + .with_device(device) + .verify(extend_lens) + .verify(seq_lens); + TensorMatcher({M, kPrefillPlanDim}) // compress_plan and write_plan + .with_dtype() + .with_device(device) + .verify(compress_plan) + .verify(write_plan); + + const auto params = CompressParams{ + .compress_plan = static_cast(compress_plan.data_ptr()), + .write_plan = static_cast(write_plan.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .extend_lens = static_cast(extend_lens.data_ptr()), + .batch_size = static_cast(N.unwrap()), + .num_tokens = static_cast(M.unwrap()), + .compress_ratio = compress_ratio, + .is_overlap = is_overlap, + }; + + if (!is_cuda) return plan_prefill_host(params, use_cuda_graph); + /// NOTE: cuda kernel plan is naturally compatible with cuda graph + LaunchKernel(1, kBlockSize, device.unwrap())(plan_prefill_cuda, params); + return PlanResult{params.num_tokens, params.num_tokens}; +} + +} // namespace host::compress + +namespace { + +[[maybe_unused]] +constexpr auto& plan_compress_prefill = host::compress::plan_prefill; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh new file mode 100644 index 0000000000..d3953578b9 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh @@ -0,0 +1,254 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include + +#include +#include + +namespace { + +using Plan = device::compress::PrefillPlan; + +/// \brief common block size for memory-bound kernel +constexpr uint32_t kBlockSize = 128; +constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + +struct FusedNormRopeParams { + void* __restrict__ input; + const void* __restrict__ weight; + float eps; + uint32_t num_works; + const void* __restrict__ handle; + const float* __restrict__ freqs_cis; + uint32_t compress_ratio; +}; + +enum class ForwardMode { + CompressExtend = 0, + CompressDecode = 1, + DefaultForward = 2, +}; + +template +__global__ void fused_norm_rope(const __grid_constant__ FusedNormRopeParams params) { + using namespace device; + using enum ForwardMode; + + constexpr int64_t kMaxVecSize = 16 / sizeof(DType); + constexpr int64_t kVecSize = std::min(kMaxVecSize, kHeadDim / kWarpThreads); + constexpr int64_t kLocalSize = kHeadDim / (kWarpThreads * kVecSize); + constexpr int64_t kRopeVecSize = kRopeDim / (kWarpThreads * 2); + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; + static_assert(kHeadDim % (kWarpThreads * kVecSize) == 0); + static_assert(kLocalSize * kVecSize * kWarpThreads == kHeadDim); + static_assert(kRopeDim % (kWarpThreads * 2) == 0); + static_assert(kRopeDim % (kVecSize * kLocalSize) == 0); + static_assert(kRopeSize <= kWarpThreads); + static_assert(kRopeVecSize == 1, "only support rope dim = 64"); + + const auto& [ + _input, _weight, eps, num_works, // norm + handle, freqs_cis, compress_ratio // rope + ] = params; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kNumWarps + warp_id; + + if (work_id >= num_works) return; + + DType* input; + int32_t position; + if constexpr (kMode == CompressExtend) { + const auto plan = static_cast(handle)[work_id]; + input = static_cast(_input) + plan.ragged_id * kHeadDim; + position = plan.position + 1 - compress_ratio; + if (plan.ragged_id == 0xFFFFFFFF) [[unlikely]] + return; + } else if constexpr (kMode == CompressDecode) { + input = static_cast(_input) + work_id * kHeadDim; + const auto seq_len = static_cast(handle)[work_id]; + if (seq_len % compress_ratio != 0) return; + position = seq_len - compress_ratio; + } else if constexpr (kMode == DefaultForward) { + input = static_cast(_input) + work_id * kHeadDim; + position = static_cast(handle)[work_id]; + } else { + static_assert(host::dependent_false_v, "Unsupported Mode"); + } + + using Storage = AlignedVector; + __shared__ Storage s_rope_input[kNumWarps][kRopeSize]; + + // prefetch freq + const auto mem_freq = tile::Memory::warp(); + const auto freq = mem_freq.load(freqs_cis + position * kRopeDim); + + PDLWaitPrimary(); + + // part 1: norm + { + const auto gmem = tile::Memory::warp(); + Storage input_vec[kLocalSize]; + Storage weight_vec[kLocalSize]; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + input_vec[i] = gmem.load(input, i); + } + +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + weight_vec[i] = gmem.load(_weight, i); + } + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const auto fp32_input = cast(input_vec[i][j]); + sum_of_squares += fp32_input * fp32_input; + } + } + + sum_of_squares = warp::reduce_sum(sum_of_squares); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + eps); + +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const auto fp32_input = cast(input_vec[i][j]); + const auto fp32_weight = cast(weight_vec[i][j]); + input_vec[i][j] = cast(fp32_input * norm_factor * fp32_weight); + } + } + + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; + +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + if (i == kLocalSize - 1 && is_rope_lane) { + const auto rope_id = lane_id - (kWarpThreads - kRopeSize); + s_rope_input[warp_id][rope_id] = input_vec[i]; + } else { + gmem.store(input, input_vec[i], i); + } + } + + __syncwarp(); + } + + // part 2: rope + { + // mem elem = DType x 2 + using DTypex2_t = packed_t; + const auto mem_elem = tile::Memory::warp(); + const auto elem = mem_elem.load(s_rope_input[warp_id]); + const auto [x_real, x_imag] = cast(elem); + const auto [freq_real, freq_imag] = freq; + const fp32x2_t output = { + x_real * freq_real - x_imag * freq_imag, + x_real * freq_imag + x_imag * freq_real, + }; + mem_elem.store(input + (kHeadDim - kRopeDim), cast(output)); + } + + PDLTriggerSecondary(); +} + +template +struct FusedNormRopeKernel { + template + static constexpr auto fused_kernel = fused_norm_rope; + + static void forward( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView weight, + const tvm::ffi::TensorView handle, + const tvm::ffi::TensorView freqs_cis, + int32_t _mode, + float eps, + uint32_t compress_ratio) { + using namespace host; + using enum ForwardMode; + + const auto mode = static_cast(_mode); + + auto B = SymbolicSize{"num_q_tokens"}; + auto N = SymbolicSize{"num_compress_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, kHeadDim}) // input + .with_dtype() + .with_device(device_) + .verify(input); + TensorMatcher({kHeadDim}) // weight + .with_dtype() + .with_device(device_) + .verify(weight); + TensorMatcher({-1, kRopeDim}) // freqs_cis + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + switch (mode) { + case CompressExtend: + TensorMatcher({N, compress::kPrefillPlanDim}) // plan + .with_dtype() + .with_device(device_) + .verify(handle); + RuntimeCheck(compress_ratio > 0); + break; + case CompressDecode: + TensorMatcher({N}) // seq_len + .with_dtype() + .with_device(device_) + .verify(handle); + RuntimeCheck(compress_ratio > 0); + break; + case DefaultForward: + TensorMatcher({N}) // position + .with_dtype() + .with_device(device_) + .verify(handle); + RuntimeCheck(compress_ratio == 0); + break; + default: + Panic("unsupported forward mode: ", static_cast(mode)); + } + + // launch kernel + const auto num_compress_tokens = static_cast(N.unwrap()); + if (num_compress_tokens == 0) return; + const auto params = FusedNormRopeParams{ + .input = input.data_ptr(), + .weight = weight.data_ptr(), + .eps = eps, + .num_works = num_compress_tokens, + .handle = handle.data_ptr(), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .compress_ratio = compress_ratio, + }; + const auto num_blocks = div_ceil(num_compress_tokens, kNumWarps); + using KernelType = std::decay_t)>; + static constexpr KernelType kernel_table[3] = { + [static_cast(CompressExtend)] = fused_kernel, + [static_cast(CompressDecode)] = fused_kernel, + [static_cast(DefaultForward)] = fused_kernel, + }; + const auto kernel = kernel_table[static_cast(mode)]; + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh new file mode 100644 index 0000000000..a9cac17544 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh @@ -0,0 +1,643 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include +#include + +#include + +#include + +namespace { + +using PlanC = device::compress::CompressPlan; +using PlanD = device::compress::DecodePlan; +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::inv_scale_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +SGL_DEVICE uint8_t quant_fp4_e2m1(float x) { + const float ax = fminf(fabsf(x), 6.0f); + uint8_t idx = 0; + idx += ax > 0.25f; + idx += ax > 0.75f; + idx += ax > 1.25f; + idx += ax > 1.75f; + idx += ax > 2.5f; + idx += ax > 3.5f; + idx += ax > 5.0f; + if (x < 0.0f && idx != 0) idx |= 0x8; + return idx; +} + +constexpr uint32_t kBlockSize = 256; +constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + +struct FusedNormRopeStoreParams { + void* __restrict__ input; + const void* __restrict__ handle; // plan decode / compress + const void* __restrict__ weight; + const float* __restrict__ freqs_cis; + const int32_t* __restrict__ out_loc; + uint8_t* __restrict__ kvcache; + float eps; + uint32_t compress_ratio; + uint32_t num_tokens; +}; + +enum class ForwardMode : bool { + CompressExtend = 0, + CompressDecode = 1, +}; + +#define INDEXER_KERNEL __global__ __launch_bounds__(kBlockSize, 8) +#define FLASHMLA_KERNEL __global__ __launch_bounds__(kBlockSize, 8) + +// ---------------------------------------------------------------------------- +// Indexer variant: kHeadDim = 128, 1 token per *warp* (8 tokens per block). +// Each warp's 32 lanes cover the full 128-elem head_dim (kVecSize = 4 each). +// Cache layout: 132 bytes/token (128 fp8 nope + 4 fp32 scale). +// ---------------------------------------------------------------------------- +template +INDEXER_KERNEL void fused_norm_rope_indexer(const __grid_constant__ FusedNormRopeStoreParams params) { + using namespace device; + using enum ForwardMode; + + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kVecSize = 4; + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; + constexpr int64_t kPageBytes = 132ll << kPageBits; + static_assert(kHeadDim == kWarpThreads * kVecSize); + static_assert(kRopeDim == kWarpThreads * 2); + static_assert(kRopeSize <= kWarpThreads); + using Storage = AlignedVector; + using Float4 = AlignedVector; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kNumWarps + warp_id; + // Lanes whose 4-elem pack lies in the rope tail (= last `kRopeSize` packs). + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; + + if (work_id >= params.num_tokens) return; + + const auto input = static_cast(params.input) + work_id * kHeadDim; + int32_t position; + int32_t out_loc; + if constexpr (kMode == CompressExtend) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.is_invalid()) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[plan.ragged_id]; + } else if constexpr (kMode == CompressDecode) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.seq_len % params.compress_ratio != 0) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[work_id]; + } else { + static_assert(host::dependent_false_v, "Unsupported Mode"); + } + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + PDLWaitPrimary(); + Float4 data, freq; + + // part 1: norm + { + Storage input_vec, weight_vec; + input_vec.load(input, lane_id); + weight_vec.load(params.weight, lane_id); + if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + sum_of_squares += fp32_input * fp32_input; + } + + sum_of_squares = warp::reduce_sum(sum_of_squares); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); + +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + const auto fp32_weight = cast(weight_vec[i]); + data[i] = fp32_input * norm_factor * fp32_weight; + } + } + + // part 2: rope (rope-lane only, 4 elems per lane = 2 (real, imag) pairs) + if (is_rope_lane) { + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto y_real = data[2]; + const auto y_imag = data[3]; + const auto freq_x_real = freq[0]; + const auto freq_x_imag = freq[1]; + const auto freq_y_real = freq[2]; + const auto freq_y_imag = freq[3]; + data[0] = x_real * freq_x_real - x_imag * freq_x_imag; + data[1] = x_real * freq_x_imag + x_imag * freq_x_real; + data[2] = y_real * freq_y_real - y_imag * freq_y_imag; + data[3] = y_real * freq_y_imag + y_imag * freq_y_real; + } + + // part 3: hadamard transform + { + // Stage 1: butterfly (data[0], data[1]) and (data[2], data[3]). + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a1; + data[1] = a0 - a1; + data[2] = a2 + a3; + data[3] = a2 - a3; + } + // Stage 2: butterfly (data[0], data[2]) and (data[1], data[3]). + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a2; + data[1] = a1 + a3; + data[2] = a0 - a2; + data[3] = a1 - a3; + } + // Stages 3..7: cross-lane butterflies. Lower-lane (mask bit clear) keeps + // the sum, upper-lane (mask bit set) keeps the difference. shfl_xor is + // unsynchronized across early-returned lanes, but invalid-plan returns + // happen above for *all* lanes of a warp (work_id is warp-uniform), so + // the warp is intact here. +#pragma unroll + for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { +#ifndef USE_ROCM + const float other = __shfl_xor_sync(kFullMask, data[i], mask, kWarpThreads); +#else + const float other = __shfl_xor(data[i], mask, kWarpThreads); +#endif + data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); + } + } + const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] *= kHadamardScale; + } + + // part 4: per-warp UE8M0 quant + store. The whole warp emits one fp8 group + // (= 128 elements) plus a single fp32 scale, matching the indexer cache + // layout (`fused_store_indexer_cache`). + { + using OutStorage = AlignedVector; + float local_max = math::abs(data[0]); +#pragma unroll + for (int i = 1; i < kVecSize; ++i) { + local_max = math::max(local_max, math::abs(data[i])); + } + const auto abs_max = warp::reduce_max(local_max); + const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto inv_scale = 1.0f / scale; + const int32_t page = out_loc >> kPageBits; + const int32_t offset = out_loc & ((1 << kPageBits) - 1); + const auto page_ptr = params.kvcache + page * kPageBytes; + const auto value_ptr = page_ptr + offset * 128; + const auto scale_ptr = page_ptr + (128 << kPageBits) + offset * 4; + OutStorage result; + result[0] = pack_fp8(data[0] * inv_scale, data[1] * inv_scale); + result[1] = pack_fp8(data[2] * inv_scale, data[3] * inv_scale); + PDLTriggerSecondary(); + result.store(value_ptr, lane_id); + // The single fp32 scale is identical across all lanes -- write from any lane. + if (lane_id == 0) reinterpret_cast(scale_ptr)[0] = scale; + } +} + +template +INDEXER_KERNEL void fused_norm_rope_indexer_fp4(const __grid_constant__ FusedNormRopeStoreParams params) { + using namespace device; + using enum ForwardMode; + + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kVecSize = 4; + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; + constexpr int64_t kPageBytes = 68ll << kPageBits; + static_assert(kHeadDim == kWarpThreads * kVecSize); + static_assert(kRopeDim == kWarpThreads * 2); + static_assert(kRopeSize <= kWarpThreads); + using Storage = AlignedVector; + using Float4 = AlignedVector; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kNumWarps + warp_id; + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; + + if (work_id >= params.num_tokens) return; + + const auto input = static_cast(params.input) + work_id * kHeadDim; + int32_t position; + int32_t out_loc; + if constexpr (kMode == CompressExtend) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.is_invalid()) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[plan.ragged_id]; + } else if constexpr (kMode == CompressDecode) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.seq_len % params.compress_ratio != 0) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[work_id]; + } else { + static_assert(host::dependent_false_v, "Unsupported Mode"); + } + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + PDLWaitPrimary(); + Float4 data, freq; + + { + Storage input_vec, weight_vec; + input_vec.load(input, lane_id); + weight_vec.load(params.weight, lane_id); + if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + sum_of_squares += fp32_input * fp32_input; + } + + sum_of_squares = warp::reduce_sum(sum_of_squares); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); + +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + const auto fp32_weight = cast(weight_vec[i]); + data[i] = fp32_input * norm_factor * fp32_weight; + } + } + + if (is_rope_lane) { + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto y_real = data[2]; + const auto y_imag = data[3]; + const auto freq_x_real = freq[0]; + const auto freq_x_imag = freq[1]; + const auto freq_y_real = freq[2]; + const auto freq_y_imag = freq[3]; + data[0] = x_real * freq_x_real - x_imag * freq_x_imag; + data[1] = x_real * freq_x_imag + x_imag * freq_x_real; + data[2] = y_real * freq_y_real - y_imag * freq_y_imag; + data[3] = y_real * freq_y_imag + y_imag * freq_y_real; + } + + { + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a1; + data[1] = a0 - a1; + data[2] = a2 + a3; + data[3] = a2 - a3; + } + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a2; + data[1] = a1 + a3; + data[2] = a0 - a2; + data[3] = a1 - a3; + } +#pragma unroll + for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); + data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); + } + } + const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] *= kHadamardScale; + } + + { + float local_max = math::abs(data[0]); +#pragma unroll + for (int i = 1; i < kVecSize; ++i) { + local_max = math::max(local_max, math::abs(data[i])); + } + local_max = warp::reduce_max<8>(local_max); + + const auto scale_raw = fmaxf(1e-4f, local_max) / 6.0f; + const auto scale_ue8m0 = static_cast(cast_to_ue8m0(scale_raw)); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + + const uint8_t packed0 = quant_fp4_e2m1(data[0] * inv_scale) | (quant_fp4_e2m1(data[1] * inv_scale) << 4); + const uint8_t packed1 = quant_fp4_e2m1(data[2] * inv_scale) | (quant_fp4_e2m1(data[3] * inv_scale) << 4); + const uint16_t packed = static_cast(packed0) | (static_cast(packed1) << 8); + + const int32_t page = out_loc >> kPageBits; + const int32_t offset = out_loc & ((1 << kPageBits) - 1); + const auto page_ptr = params.kvcache + page * kPageBytes; + const auto value_ptr = page_ptr + offset * 64; + const auto scale_ptr = page_ptr + (64 << kPageBits) + offset * 4; + + PDLTriggerSecondary(); + reinterpret_cast(value_ptr)[lane_id] = packed; + if ((lane_id & 7) == 0) static_cast(scale_ptr)[lane_id >> 3] = scale_ue8m0; + } +} + +// ---------------------------------------------------------------------------- +// FlashMLA variant: kHeadDim = 512, 1 token per *block* (256 threads). +// Each thread loads kVecSize=2 BF16, so 256 threads cover the full 512 elems. +// Cache layout: 584 bytes/token = 448 fp8 nope + 64 (=32 bf16x2) rope + 8 scale. +// ---------------------------------------------------------------------------- +template +FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormRopeStoreParams params) { + using namespace device; + using enum ForwardMode; + + constexpr int64_t kHeadDim = 512; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kVecSize = 2; + // Last warp owns the rope tail. The remaining 7 warps each emit one + // 64-element fp8 group (own UE8M0 scale). + constexpr uint32_t kRopeWarp = kNumWarps - 1; + constexpr int64_t kPageBytes = host::div_ceil(584ll << kPageBits, 576) * 576; + static_assert(kHeadDim == kBlockSize * kVecSize); + static_assert(kRopeDim == kWarpThreads * kVecSize); + static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); + using Storage = AlignedVector; + using Float2 = AlignedVector; + + const auto tx = threadIdx.x; + const auto warp_id = tx / kWarpThreads; + const auto lane_id = tx % kWarpThreads; + const auto work_id = blockIdx.x; + + if (work_id >= params.num_tokens) return; + + const auto input = static_cast(params.input) + work_id * kHeadDim; + int32_t position; + int32_t out_loc; + if constexpr (kMode == CompressExtend) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.is_invalid()) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[plan.ragged_id]; + } else if constexpr (kMode == CompressDecode) { + const auto plan = static_cast(params.handle)[work_id]; + if (plan.seq_len % params.compress_ratio != 0) return; + position = plan.seq_len - params.compress_ratio; + out_loc = params.out_loc[work_id]; + } else { + static_assert(host::dependent_false_v, "Unsupported Mode"); + } + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + PDLWaitPrimary(); + Float2 data, freq; + + // part 1: norm. Each thread owns one 2-elem pack (`tx`-th pack of input). + // Sum of squares is reduced across the whole block via per-warp partials. + { + __shared__ float partial_sums[kNumWarps]; + + Storage input_vec, weight_vec; + input_vec.load(input, tx); + weight_vec.load(params.weight, tx); + if (warp_id == kRopeWarp) freq.load(freqs_cis, lane_id); + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + sum_of_squares += fp32_input * fp32_input; + } + + const auto warp_sum = warp::reduce_sum(sum_of_squares); + if (lane_id == 0) partial_sums[warp_id] = warp_sum; + __syncthreads(); + // Replicate the per-warp partial sums to a full warp and reduce. Every + // lane-group of `kNumWarps` lanes ends up with the global sum. + sum_of_squares = warp::reduce_sum(partial_sums[lane_id % kNumWarps]); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); + +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto fp32_input = cast(input_vec[i]); + const auto fp32_weight = cast(weight_vec[i]); + data[i] = fp32_input * norm_factor * fp32_weight; + } + } + + const int32_t page = out_loc >> kPageBits; + const int32_t offset = out_loc & ((1 << kPageBits) - 1); + const auto page_ptr = params.kvcache + page * kPageBytes; + const auto value_ptr = page_ptr + offset * 576; + + PDLTriggerSecondary(); + + // part 2: rope on the rope warp (BF16 store), or per-warp FP8 quant + store. + if (warp_id == kRopeWarp) { + // Each rope-warp lane owns exactly one (real, imag) pair within the rope + // tail. Apply rotation, downcast to BF16, write to the slot's rope region. + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto freq_real = freq[0]; + const auto freq_imag = freq[1]; + data[0] = x_real * freq_real - x_imag * freq_imag; + data[1] = x_real * freq_imag + x_imag * freq_real; + const auto result = cast(fp32x2_t{data[0], data[1]}); + const auto rope_ptr = value_ptr + 448; + reinterpret_cast(rope_ptr)[lane_id] = result; + } else { + // Non-rope warp: per-warp UE8M0 group (64 elems -> 64 fp8 + 1 scale byte). + // BF16 round-trip to match the precision of the non-fused path + // (which goes through quant_to_nope_fp8_rope_bf16_pack_triton with bf16 input). + const auto x = cast(cast(data[0])); + const auto y = cast(cast(data[1])); + const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); + const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + const auto result = pack_fp8(x * inv_scale, y * inv_scale); + const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; + reinterpret_cast(value_ptr)[tx] = result; + // All lanes in this warp produce the same scale byte; let lane 0 publish. + if (lane_id == 0) static_cast(scale_ptr)[warp_id] = scale_ue8m0; + } +} + +template +struct FusedNormRopeKernel { + static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); + static constexpr bool kIsIndexer = (kHeadDim == 128); + static constexpr int64_t kIndexerBytes = 132 * kPageSize; + static constexpr int64_t kFlashMLABytes = host::div_ceil(584 * kPageSize, 576) * 576; + static constexpr int64_t kPageBytes = kIsIndexer ? kIndexerBytes : kFlashMLABytes; + + /// TODO: Let's fix the config for now. + static_assert(kRopeDim == 64 && (kHeadDim == 128 || kHeadDim == 512)); + static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); + + template + static constexpr auto select_kernel() { + if constexpr (kIsIndexer) { + return fused_norm_rope_indexer; + } else { + return fused_norm_rope_flashmla; + } + } + + template + static constexpr auto select_fp4_kernel() { + static_assert(kIsIndexer, "FP4 fused store is only defined for the indexer"); + return fused_norm_rope_indexer_fp4; + } + + static void forward( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView plan, + const tvm::ffi::TensorView weight, + const float eps, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView out_loc, + const tvm::ffi::TensorView kvcache, + const bool is_decode, + const uint32_t compress_ratio) { + using namespace host; + using enum ForwardMode; + + const auto mode = static_cast(is_decode); + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, kHeadDim}) // input + .with_dtype() + .with_device(device_) + .verify(input); + TensorMatcher({kHeadDim}) // weight + .with_dtype() + .with_device(device_) + .verify(weight); + TensorMatcher({-1, kRopeDim}) // freqs_cis + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + TensorMatcher({-1}) // out_loc + .with_dtype() + .with_device(device_) + .verify(out_loc); + TensorMatcher({-1, -1}) // cache + .with_strides({kPageBytes, 1}) + .with_dtype() + .with_device(device_) + .verify(kvcache); + + switch (mode) { + case CompressExtend: + compress::verify_plan_c(plan, N, device_); + RuntimeCheck(out_loc.size(0) >= N.unwrap()); + break; + case CompressDecode: + compress::verify_plan_d(plan, N, device_); + RuntimeCheck(out_loc.size(0) == N.unwrap()); + break; + } + + const auto num_tokens = static_cast(N.unwrap()); + if (num_tokens == 0) return; + const auto params = FusedNormRopeStoreParams{ + .input = input.data_ptr(), + .handle = plan.data_ptr(), + .weight = weight.data_ptr(), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .out_loc = static_cast(out_loc.data_ptr()), + .kvcache = static_cast(kvcache.data_ptr()), + .eps = eps, + .compress_ratio = compress_ratio, + .num_tokens = num_tokens, + }; + // Indexer packs `kNumWarps` tokens per block (warp-major); FlashMLA uses + // a whole block per token (cta-major sum-reduce over head_dim=512). + const uint32_t num_blocks = kIsIndexer ? div_ceil(num_tokens, kNumWarps) : num_tokens; + const auto device = device_.unwrap(); + const auto kernel = mode == CompressExtend ? select_kernel() : select_kernel(); + LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params); + } + + static void forward_fp4( + const tvm::ffi::TensorView input, + const tvm::ffi::TensorView plan, + const tvm::ffi::TensorView weight, + const float eps, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView out_loc, + const tvm::ffi::TensorView kvcache, + const bool is_decode, + const uint32_t compress_ratio) { + using namespace host; + using enum ForwardMode; + + static_assert(kIsIndexer, "FP4 fused store is only defined for the indexer"); + constexpr int64_t kFp4PageBytes = 68 * kPageSize; + const auto mode = static_cast(is_decode); + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({N, kHeadDim}).with_dtype().with_device(device_).verify(input); + TensorMatcher({kHeadDim}).with_dtype().with_device(device_).verify(weight); + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + TensorMatcher({-1}).with_dtype().with_device(device_).verify(out_loc); + TensorMatcher({-1, -1}).with_strides({kFp4PageBytes, 1}).with_dtype().with_device(device_).verify(kvcache); + + switch (mode) { + case CompressExtend: + compress::verify_plan_c(plan, N, device_); + RuntimeCheck(out_loc.size(0) >= N.unwrap()); + break; + case CompressDecode: + compress::verify_plan_d(plan, N, device_); + RuntimeCheck(out_loc.size(0) == N.unwrap()); + break; + } + + const auto num_tokens = static_cast(N.unwrap()); + if (num_tokens == 0) return; + const auto params = FusedNormRopeStoreParams{ + .input = input.data_ptr(), + .handle = plan.data_ptr(), + .weight = weight.data_ptr(), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .out_loc = static_cast(out_loc.data_ptr()), + .kvcache = static_cast(kvcache.data_ptr()), + .eps = eps, + .compress_ratio = compress_ratio, + .num_tokens = num_tokens, + }; + const uint32_t num_blocks = div_ceil(num_tokens, kNumWarps); + const auto device = device_.unwrap(); + const auto kernel = + mode == CompressExtend ? select_fp4_kernel() : select_fp4_kernel(); + LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh new file mode 100644 index 0000000000..90dec3c117 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh @@ -0,0 +1,214 @@ +#include +#include + +#include +#include +#include + +#include + +#include +#include + +namespace { + +[[maybe_unused]] +SGL_DEVICE float act_sqrt_softplus(float x) { + const float softplus = fmaxf(x, 0.0f) + log1pf(expf(-fabsf(x))); + return sqrtf(softplus); +} + +struct MoEHashTopKParams { + const float* __restrict__ router_logits; + const int64_t* __restrict__ input_id; + const int32_t* __restrict__ tid2eid; + int32_t* __restrict__ topk_ids; + float* __restrict__ topk_weights; + uint32_t num_tokens; + uint32_t topk; + uint32_t num_routed_experts; + uint32_t num_shared_experts; + float routed_scaling_factor; +}; + +template +__global__ void moe_hash_topk_fused(const MoEHashTopKParams __grid_constant__ params) { + using namespace device; + const auto& [ + router_logits, input_id, tid2eid, topk_ids, topk_weights, // pointers + num_tokens, topk, num_routed_experts, num_shared_experts, routed_scaling_factor] = + params; + + const uint32_t topk_fused = topk + num_shared_experts; + const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t warp_id = tid / kWarpThreads; + const uint32_t lane_id = tid % kWarpThreads; + if (warp_id >= num_tokens) return; + // we can safely prefetch the token id + const auto token_id = input_id[warp_id]; + + PDLWaitPrimary(); + + float routed_weight = 0.0f; + int32_t expert_id = 0; + if (lane_id < topk) { + expert_id = tid2eid[token_id * topk + lane_id]; + routed_weight = Fn(router_logits[warp_id * num_routed_experts + expert_id]); + } + + const auto routed_sum = device::warp::reduce_sum(routed_weight); + if (lane_id < topk_fused) { + const bool is_shared = lane_id >= topk; + const auto output_offset = warp_id * topk_fused + lane_id; + topk_ids[output_offset] = is_shared ? num_routed_experts + lane_id - topk : expert_id; + topk_weights[output_offset] = is_shared ? 1.0f / routed_scaling_factor : routed_weight / routed_sum; + } + + PDLTriggerSecondary(); +} + +struct TopKParams { + int32_t* __restrict__ topk_ids; + // Exactly one is active: ntn_ptr == nullptr means use ntn_value. + const int32_t* __restrict__ ntn_ptr; + int32_t ntn_value; + int64_t stride; + uint32_t topk; + uint32_t num_tokens; +}; + +__global__ void mask_topk_ids_padded_region(const TopKParams __grid_constant__ params) { + const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; + const uint32_t warp_id = tid / device::kWarpThreads; + const uint32_t lane_id = tid % device::kWarpThreads; + if (warp_id >= params.num_tokens || lane_id >= params.topk) return; + device::PDLWaitPrimary(); + const uint32_t num = (params.ntn_ptr != nullptr) // + ? static_cast(params.ntn_ptr[0]) + : static_cast(params.ntn_value); + if (warp_id >= num) params.topk_ids[warp_id * params.stride + lane_id] = -1; + device::PDLTriggerSecondary(); +} + +template +struct HashTopKKernel { + static constexpr auto kernel = moe_hash_topk_fused; + + static void + run(const tvm::ffi::TensorView router_logits, + const tvm::ffi::TensorView input_id, + const tvm::ffi::TensorView tid2eid, + const tvm::ffi::TensorView topk_weights, + const tvm::ffi::TensorView topk_ids, + float routed_scaling_factor) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto E = SymbolicSize{"num_routed_experts"}; + auto K = SymbolicSize{"topk_fused"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({N, E}) // + .with_dtype() + .with_device(device) + .verify(router_logits); + TensorMatcher({N}) // + .with_dtype() + .with_device(device) + .verify(input_id); + TensorMatcher({-1, -1}) // + .with_dtype() + .with_device(device) + .verify(tid2eid); + TensorMatcher({N, K}) // + .with_dtype() + .with_device(device) + .verify(topk_weights); + TensorMatcher({N, K}) // + .with_dtype() + .with_device(device) + .verify(topk_ids); + + const auto num_tokens = static_cast(N.unwrap()); + const auto topk_fused = static_cast(K.unwrap()); + const auto topk = static_cast(tid2eid.size(1)); + const auto shared_experts = topk_fused - topk; + RuntimeCheck(topk <= topk_fused, "HashTopKKernel requires topk <= topk_fused"); + RuntimeCheck(topk_fused <= device::kWarpThreads, "HashTopKKernel requires topk_fused <= warp size"); + + const auto params = MoEHashTopKParams{ + .router_logits = static_cast(router_logits.data_ptr()), + .input_id = static_cast(input_id.data_ptr()), + .tid2eid = static_cast(tid2eid.data_ptr()), + .topk_ids = static_cast(topk_ids.data_ptr()), + .topk_weights = static_cast(topk_weights.data_ptr()), + .num_tokens = num_tokens, + .topk = topk, + .num_routed_experts = static_cast(E.unwrap()), + .num_shared_experts = shared_experts, + .routed_scaling_factor = routed_scaling_factor, + }; + const auto kBlockSize = 128u; + const auto kNumWarps = kBlockSize / device::kWarpThreads; + const auto num_blocks = div_ceil(num_tokens, kNumWarps); + LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +// TODO this may not be related to *hash* topk, thus may move +struct MaskKernel { + static constexpr auto kernel = mask_topk_ids_padded_region; + + static void run(tvm::ffi::TensorView topk_ids, tvm::ffi::TensorView num_token_non_padded) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto K = SymbolicSize{"topk"}; + auto D = SymbolicSize{"stride"}; + auto device = SymbolicDevice{}; + device.set_options(); + TensorMatcher({N, K}) // + .with_strides({D, 1}) + .with_dtype() + .with_device(device) + .verify(topk_ids); + RuntimeCheck(num_token_non_padded.numel() == 1, "num_token_non_padded should be a scalar"); + RuntimeCheck(K.unwrap() <= device::kWarpThreads, "MaskKernel requires topk <= warp size"); + const int32_t* ntn_ptr = nullptr; + int32_t ntn_value = 0; + const auto ntn_dev = num_token_non_padded.device().device_type; + if (ntn_dev == kDLCUDA) { + RuntimeCheck(is_type(num_token_non_padded.dtype()), "num_token_non_padded on CUDA must be int32"); + ntn_ptr = static_cast(num_token_non_padded.data_ptr()); + } else if (ntn_dev == kDLCPU) { + if (is_type(num_token_non_padded.dtype())) { + ntn_value = *static_cast(num_token_non_padded.data_ptr()); + } else if (is_type(num_token_non_padded.dtype())) { + ntn_value = static_cast(*static_cast(num_token_non_padded.data_ptr())); + } else { + RuntimeCheck(false, "num_token_non_padded on CPU must be int32 or int64"); + } + } else { + RuntimeCheck(false, "num_token_non_padded must be on CPU or CUDA"); + } + + const auto num_tokens = static_cast(N.unwrap()); + const auto params = TopKParams{ + .topk_ids = static_cast(topk_ids.data_ptr()), + .ntn_ptr = ntn_ptr, + .ntn_value = ntn_value, + .stride = static_cast(D.unwrap()), + .topk = static_cast(K.unwrap()), + .num_tokens = num_tokens, + }; + const auto kBlockSize = 128u; + const auto kNumWarps = kBlockSize / device::kWarpThreads; + const auto num_blocks = div_ceil(num_tokens, kNumWarps); + LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // + .enable_pdl(true)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh new file mode 100644 index 0000000000..aefec24372 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh @@ -0,0 +1,82 @@ +#include +#include + +#include + +#include + +#include +#include + +#include + +namespace { + +/// NOTE: for offload to cpu kernel, we use persistent kernel +inline constexpr uint32_t kBlockSize = 1024; +inline constexpr uint32_t kBlockQuota = 4; + +#define OFFLOAD_KERNEL __global__ __launch_bounds__(kBlockSize, 1) + +struct OffloadParams { + void** gpu_caches; + void** cpu_caches; + const int64_t* gpu_indices; + const int64_t* cpu_indices; + uint32_t num_items; + uint32_t num_layers; +}; + +OFFLOAD_KERNEL void offload_to_cpu(const __grid_constant__ OffloadParams params) { + using namespace device::hisparse; + const auto [gpu_caches, cpu_caches, gpu_indices, cpu_indices, num_items, num_layers] = params; + const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; + constexpr auto kNumWarps = (kBlockSize / 32) * kBlockQuota; + for (auto i = global_tid / 32; i < num_items; i += kNumWarps) { + const int32_t gpu_index = gpu_indices[i]; + const int32_t cpu_index = cpu_indices[i]; + for (auto j = 0u; j < num_layers; ++j) { + const auto gpu_cache = gpu_caches[j]; + const auto cpu_cache = cpu_caches[j]; + transfer_item( + /*dst_cache=*/cpu_cache, + /*src_cache=*/gpu_cache, + /*dst_index=*/cpu_index, + /*src_index=*/gpu_index); + } + } +} + +[[maybe_unused]] +void hisparse_transfer( + tvm::ffi::TensorView gpu_ptrs, + tvm::ffi::TensorView cpu_ptrs, + tvm::ffi::TensorView gpu_indices, + tvm::ffi::TensorView cpu_indices) { + using namespace host; + auto N = SymbolicSize{"num_items"}; + auto L = SymbolicSize{"num_layers"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({L}) // 1D cache pointers + .with_dtype() + .with_device(device_) + .verify(gpu_ptrs) + .verify(cpu_ptrs); + TensorMatcher({N}) // 1D indices + .with_dtype() + .with_device(device_) + .verify(gpu_indices) + .verify(cpu_indices); + const auto params = OffloadParams{ + .gpu_caches = static_cast(gpu_ptrs.data_ptr()), + .cpu_caches = static_cast(cpu_ptrs.data_ptr()), + .gpu_indices = static_cast(gpu_indices.data_ptr()), + .cpu_indices = static_cast(cpu_indices.data_ptr()), + .num_items = static_cast(N.unwrap()), + .num_layers = static_cast(L.unwrap()), + }; + LaunchKernel(kBlockQuota, kBlockSize, device_.unwrap())(offload_to_cpu, params); +} + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh new file mode 100644 index 0000000000..8fc8d0821d --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh @@ -0,0 +1,845 @@ +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include + +#include +#include + +namespace { + +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::inv_scale_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +SGL_DEVICE uint8_t quant_fp4_e2m1(float x) { + const float ax = fminf(fabsf(x), 6.0f); + uint8_t idx = 0; + idx += ax > 0.25f; + idx += ax > 0.75f; + idx += ax > 1.25f; + idx += ax > 1.75f; + idx += ax > 2.5f; + idx += ax > 3.5f; + idx += ax > 5.0f; + if (x < 0.0f && idx != 0) idx |= 0x8; + return idx; +} + +// 4 warps per block: warp-per-(token, head) work-item dispatch (Q kernel). +constexpr uint32_t kFusedQBlockSize = 128; +constexpr uint32_t kFusedQNumWarps = kFusedQBlockSize / device::kWarpThreads; + +// 8 warps per block: block-per-token work-item dispatch (K kernel). +constexpr uint32_t kFusedKBlockSize = 256; +constexpr uint32_t kFusedKNumWarps = kFusedKBlockSize / device::kWarpThreads; + +#define Q_KERNEL __global__ __launch_bounds__(kFusedQBlockSize, 16) +#define K_KERNEL __global__ __launch_bounds__(kFusedKBlockSize, 8) + +// ============================================================================ +// Q kernel: warp-per-(token, head) rmsnorm-self + RoPE + write to q_out. +// ============================================================================ + +struct FusedQNormRopeParams { + const void* __restrict__ q_input; // (B, num_q_heads, kHeadDim) DType + void* __restrict__ q_output; // (B, num_q_heads, kHeadDim) DType + const float* __restrict__ freqs_cis; // (max_pos, kRopeDim) fp32 (re/im interleaved) + const void* __restrict__ positions; // (B,) PosT + int64_t q_input_stride_batch; + int64_t q_output_stride_batch; + uint32_t batch_size; + uint32_t num_q_heads; + float eps; +}; + +template +Q_KERNEL void fused_q_norm_rope(const __grid_constant__ FusedQNormRopeParams params) { + using namespace device; + + constexpr int64_t kMaxVecSize = 16 / sizeof(DType); + constexpr int64_t kVecSize = std::min(kMaxVecSize, kHeadDim / kWarpThreads); + constexpr int64_t kLocalSize = kHeadDim / (kWarpThreads * kVecSize); + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; + static_assert(kHeadDim % (kWarpThreads * kVecSize) == 0); + static_assert(kLocalSize * kVecSize * kWarpThreads == kHeadDim); + static_assert(kRopeDim % kVecSize == 0); + static_assert(kRopeSize <= kWarpThreads); + static_assert(kRopeDim == kWarpThreads * 2, "1 (real, imag) pair per lane"); + + using Storage = AlignedVector; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; + + const uint32_t total_works = params.batch_size * params.num_q_heads; + if (work_id >= total_works) return; + + const uint32_t batch_id = work_id / params.num_q_heads; + const uint32_t head_id = work_id % params.num_q_heads; + const auto input_ptr = + static_cast(params.q_input) + batch_id * params.q_input_stride_batch + head_id * kHeadDim; + const auto output_ptr = + static_cast(params.q_output) + batch_id * params.q_output_stride_batch + head_id * kHeadDim; + const auto position = static_cast(static_cast(params.positions)[batch_id]); + + __shared__ Storage s_rope[kFusedQNumWarps][kRopeSize]; + + // Prefetch this lane's freq pair before the PDL gate so the wait happens + // outside the dependency chain on `position`. + const auto mem_freq = tile::Memory{lane_id, kWarpThreads}; + + PDLWaitPrimary(); + + // part 1: rmsnorm-self (no weight). + const auto gmem = tile::Memory{lane_id, kWarpThreads}; + Storage input_vec[kLocalSize]; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + input_vec[i] = gmem.load(input_ptr, i); + } + + const auto freq = mem_freq.load(params.freqs_cis + position * kRopeDim); + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const auto x = cast(input_vec[i][j]); + sum_of_squares += x * x; + } + } + sum_of_squares = warp::reduce_sum(sum_of_squares); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); + +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const auto x = cast(input_vec[i][j]); + input_vec[i][j] = cast(x * norm_factor); + } + } + + // Stash the rope tail (last kRopeSize lanes' last tile) into shared memory; + // write nope tiles to gmem directly. + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + if (i == kLocalSize - 1 && is_rope_lane) { + const auto rope_id = lane_id - (kWarpThreads - kRopeSize); + s_rope[warp_id][rope_id] = input_vec[i]; + } else { + gmem.store(output_ptr, input_vec[i], i); + } + } + __syncwarp(); + + PDLTriggerSecondary(); + + // part 2: RoPE on all 32 lanes -- one (real, imag) bf16x2 pair per lane. + using DType2 = packed_t; + const auto mem_elem = tile::Memory{lane_id, kWarpThreads}; + const auto elem = mem_elem.load(s_rope[warp_id]); + const auto [x_real, x_imag] = cast(elem); + const auto [freq_real, freq_imag] = freq; + const fp32x2_t rotated = { + x_real * freq_real - x_imag * freq_imag, + x_real * freq_imag + x_imag * freq_real, + }; + mem_elem.store(output_ptr + (kHeadDim - kRopeDim), cast(rotated)); +} + +template +struct FusedQNormRopeKernel { + template + static constexpr auto kernel = fused_q_norm_rope; + + static void forward( + const tvm::ffi::TensorView q_input, + const tvm::ffi::TensorView q_output, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + float eps) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto H = SymbolicSize{"num_q_heads"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, H, kHeadDim}) // + .with_strides({-1, kHeadDim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_input); + TensorMatcher({B, H, kHeadDim}) // + .with_strides({-1, kHeadDim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_output); + TensorMatcher({-1, kRopeDim}) // + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({B}) // + .with_dtype(pos_dtype) + .with_device(device_) + .verify(positions); + + const auto batch_size = static_cast(B.unwrap()); + const auto num_q_heads = static_cast(H.unwrap()); + if (batch_size == 0) return; + + const auto params = FusedQNormRopeParams{ + .q_input = q_input.data_ptr(), + .q_output = q_output.data_ptr(), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .q_input_stride_batch = q_input.stride(0), + .q_output_stride_batch = q_output.stride(0), + .batch_size = batch_size, + .num_q_heads = num_q_heads, + .eps = eps, + }; + const auto total_works = batch_size * num_q_heads; + const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); + const auto k_int32 = kernel; + const auto k_int64 = kernel; + const auto k = pos_dtype.is_type() ? k_int32 : k_int64; + LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +// ============================================================================ +// K kernel: block-per-token rmsnorm (with kv_weight) + RoPE + FlashMLA store. +// ============================================================================ + +struct FusedKNormRopeFlashMLAParams { + const void* __restrict__ kv; // (B, kHeadDim) DType + const void* __restrict__ kv_weight; // (kHeadDim,) DType + const float* __restrict__ freqs_cis; // (max_pos, kRopeDim) fp32 + const void* __restrict__ positions; // (B,) PosT + const int32_t* __restrict__ out_loc; // (B,) int32 -> cache slot id + uint8_t* __restrict__ kvcache; // (npages, kPageBytes) uint8 + // Row stride for `kv` in elements. Required because the upstream caller often + // passes `qkv_a[..., q_lora_rank:]`, a non-contiguous slice whose stride[0] + // equals `q_lora_rank + kHeadDim` rather than `kHeadDim`. + int64_t kv_stride_batch; + uint32_t batch_size; + float eps; +}; + +template +K_KERNEL void fused_k_norm_rope_flashmla(const __grid_constant__ FusedKNormRopeFlashMLAParams params) { + using namespace device; + + constexpr int64_t kVecSize = 2; + constexpr uint32_t kRopeWarp = kFusedKNumWarps - 1; + constexpr int64_t kPageBytes = host::div_ceil(584ll << kPageBits, 576) * 576; + static_assert(kHeadDim == kFusedKBlockSize * kVecSize); + static_assert(kRopeDim == kWarpThreads * kVecSize); + static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); + using Storage = AlignedVector; + using Float2 = AlignedVector; + + const auto tx = threadIdx.x; + const auto warp_id = tx / kWarpThreads; + const auto lane_id = tx % kWarpThreads; + const auto work_id = blockIdx.x; + if (work_id >= params.batch_size) return; + + const auto input_ptr = static_cast(params.kv) + work_id * params.kv_stride_batch; + const auto position = static_cast(static_cast(params.positions)[work_id]); + const auto out_loc = params.out_loc[work_id]; + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + PDLWaitPrimary(); + Float2 data, freq; + + // part 1: norm. Each thread owns one 2-elem pack (the `tx`-th). + // Sum-of-squares is reduced block-wide via per-warp partials. + { + __shared__ float partial_sums[kFusedKNumWarps]; + + Storage input_vec, weight_vec; + input_vec.load(input_ptr, tx); + weight_vec.load(params.kv_weight, tx); + if (warp_id == kRopeWarp) freq.load(freqs_cis, lane_id); + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto x = cast(input_vec[i]); + sum_of_squares += x * x; + } + const auto warp_sum = warp::reduce_sum(sum_of_squares); + if (lane_id == 0) partial_sums[warp_id] = warp_sum; + __syncthreads(); + // Replicate the per-warp partial sums onto all lanes of one warp and + // reduce. Every group of `kBlockItemNumWarps` lanes ends up with the + // global sum. + sum_of_squares = warp::reduce_sum(partial_sums[lane_id % kFusedKNumWarps]); + const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); + +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const auto x = cast(input_vec[i]); + const auto w = cast(weight_vec[i]); + data[i] = x * norm_factor * w; + } + } + + const int32_t page = out_loc >> kPageBits; + const int32_t offset = out_loc & ((1 << kPageBits) - 1); + const auto page_ptr = params.kvcache + page * kPageBytes; + const auto value_ptr = page_ptr + offset * 576; + + PDLTriggerSecondary(); + + // part 2: rope on warp 7 (BF16 store), per-warp UE8M0 quant + store on warps 0..6. + if (warp_id == kRopeWarp) { + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto freq_real = freq[0]; + const auto freq_imag = freq[1]; + data[0] = x_real * freq_real - x_imag * freq_imag; + data[1] = x_real * freq_imag + x_imag * freq_real; + const auto result = cast(fp32x2_t{data[0], data[1]}); + const auto rope_ptr = value_ptr + 448; + reinterpret_cast(rope_ptr)[lane_id] = result; + } else { + const auto x = data[0]; + const auto y = data[1]; + const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); + const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + const auto result = pack_fp8(x * inv_scale, y * inv_scale); + const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; + reinterpret_cast(value_ptr)[tx] = result; + if (lane_id == 0) static_cast(scale_ptr)[warp_id] = scale_ue8m0; + } +} + +template +struct FusedKNormRopeFlashMLAKernel { + static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); + static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; + static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); + static_assert(1 << kLogPageSize == kPageSize); + static_assert(kHeadDim == 512 && kRopeDim == 64, "FlashMLA layout requires (512, 64)"); + + template + static constexpr auto kernel = fused_k_norm_rope_flashmla; + + static void forward( + const tvm::ffi::TensorView kv, + const tvm::ffi::TensorView kv_weight, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + const tvm::ffi::TensorView out_loc, + const tvm::ffi::TensorView kvcache, + float eps) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, kHeadDim}) // + .with_strides({-1, 1}) + .with_dtype() + .with_device(device_) + .verify(kv); + TensorMatcher({kHeadDim}) // + .with_dtype() + .with_device(device_) + .verify(kv_weight); + TensorMatcher({-1, kRopeDim}) // + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({B}) // + .with_dtype(pos_dtype) + .with_device(device_) + .verify(positions); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(out_loc); + TensorMatcher({-1, -1}) // + .with_strides({kPageBytes, 1}) + .with_dtype() + .with_device(device_) + .verify(kvcache); + + const auto batch_size = static_cast(B.unwrap()); + if (batch_size == 0) return; + + const auto params = FusedKNormRopeFlashMLAParams{ + .kv = kv.data_ptr(), + .kv_weight = kv_weight.data_ptr(), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .out_loc = static_cast(out_loc.data_ptr()), + .kvcache = static_cast(kvcache.data_ptr()), + .kv_stride_batch = kv.stride(0), + .batch_size = batch_size, + .eps = eps, + }; + const auto k_int32 = kernel; + const auto k_int64 = kernel; + const auto k = pos_dtype.is_type() ? k_int32 : k_int64; + LaunchKernel(batch_size, kFusedKBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +// ============================================================================ +// Indexer Q kernel: warp-per-(token, head) RoPE + Hadamard + fp8 act-quant. +// ============================================================================ + +struct FusedQIndexerRopeHadamardQuantParams { + const void* __restrict__ q_input; // (B, num_heads, 128) DType + void* __restrict__ q_fp8; // (B, num_heads, 128) fp8_e4m3 + // weights_out[b, h] = weight[b, h] * weight_scale * q_scale[b, h]. + // q_scale is computed internally and not exposed -- the only consumer of + // it is `weights_out`. + const void* __restrict__ weight; // (B, num_heads) DType + float* __restrict__ weights_out; // (B, num_heads) fp32 (== (B, H, 1) flat) + float weight_scale; // scalar c4_indexer.weight_scale + const float* __restrict__ freqs_cis; // (max_pos, 64) fp32 + const void* __restrict__ positions; // (B,) PosT + uint32_t batch_size; + uint32_t num_heads; +}; + +template +Q_KERNEL void fused_q_indexer_rope_hadamard_quant(const __grid_constant__ FusedQIndexerRopeHadamardQuantParams params) { + using namespace device; + + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kVecSize = 4; + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; // = 16 + static_assert(kHeadDim == kWarpThreads * kVecSize); + static_assert(kRopeDim == kWarpThreads * 2); + static_assert(kRopeSize <= kWarpThreads); + + using Storage = AlignedVector; + using Float4 = AlignedVector; + using OutStorage = AlignedVector; // 4 fp8 / lane + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; + // Last `kRopeSize` lanes own the rope tail; their 4-elem packs cover the + // trailing kRopeDim elements. + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; + + const uint32_t total_works = params.batch_size * params.num_heads; + if (work_id >= total_works) return; + + const uint32_t batch_id = work_id / params.num_heads; + const auto input_ptr = static_cast(params.q_input) + work_id * kHeadDim; + const auto position = static_cast(static_cast(params.positions)[batch_id]); + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + // Lane 0 prefetches the weight scalar for this (token, head) work item. + // Weight is (B, num_heads) DType; we need one scalar per warp -- offload + // the load to lane 0 only. The multiply + store happens once the q_scale + // is known (part 4). + + PDLWaitPrimary(); + Float4 data, freq; + const auto weight_val = cast(static_cast(params.weight)[work_id]); + + // part 1: load (no norm). Each lane owns a 4-elem pack. + { + Storage input_vec; + input_vec.load(input_ptr, lane_id); + if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + data[i] = cast(input_vec[i]); + } + } + + // part 2: rope on rope lanes only (4 elems / lane = 2 (real, imag) pairs). + if (is_rope_lane) { + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto y_real = data[2]; + const auto y_imag = data[3]; + const auto fxr = freq[0]; + const auto fxi = freq[1]; + const auto fyr = freq[2]; + const auto fyi = freq[3]; + data[0] = x_real * fxr - x_imag * fxi; + data[1] = x_real * fxi + x_imag * fxr; + data[2] = y_real * fyr - y_imag * fyi; + data[3] = y_real * fyi + y_imag * fyr; + } + + PDLTriggerSecondary(); + + // part 3: 128-point Hadamard (2 local stages + 5 cross-lane shfl_xor stages). + // Same recipe as `fused_norm_rope_indexer`; see comments there for the + // butterfly invariants and the early-return safety argument. + { + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a1; + data[1] = a0 - a1; + data[2] = a2 + a3; + data[3] = a2 - a3; + } + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a2; + data[1] = a1 + a3; + data[2] = a0 - a2; + data[3] = a1 - a3; + } +#pragma unroll + for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); + data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); + } + } + const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] *= kHadamardScale; + } + + { + float local_max = math::abs(data[0]); +#pragma unroll + for (int i = 1; i < kVecSize; ++i) { + local_max = math::max(local_max, math::abs(data[i])); + } + const auto abs_max = warp::reduce_max(local_max); + const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto inv_scale = 1.0f / scale; + OutStorage result; + result[0] = pack_fp8(data[0] * inv_scale, data[1] * inv_scale); + result[1] = pack_fp8(data[2] * inv_scale, data[3] * inv_scale); + + // q_fp8 row pointer: 128 fp8 / row = 32 OutStorage / row, one per lane. + auto out_row = static_cast(params.q_fp8) + work_id * kHeadDim; + result.store(out_row, lane_id); + params.weights_out[work_id] = weight_val * params.weight_scale * scale; + } +} + +template +struct FusedQIndexerRopeHadamardQuantKernel { + template + static constexpr auto kernel = fused_q_indexer_rope_hadamard_quant; + + static void forward( + const tvm::ffi::TensorView q_input, + const tvm::ffi::TensorView q_fp8, + const tvm::ffi::TensorView weight, + const tvm::ffi::TensorView weights_out, + double weight_scale, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions) { + using namespace host; + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + + auto B = SymbolicSize{"batch_size"}; + auto H = SymbolicSize{"num_heads"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + // Caller path is `wq_b(q_lora).view(-1, H, D)` -> contiguous; the kernel + // assumes a flat `(B*H, kHeadDim)` layout for both q_input and q_fp8. + // Pin the head/innermost strides; assert the batch stride below. + TensorMatcher({B, H, kHeadDim}) // + .with_strides({-1, kHeadDim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_input); + TensorMatcher({B, H, kHeadDim}) // + .with_strides({-1, kHeadDim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_fp8); + TensorMatcher({B, H}) // + .with_dtype() + .with_device(device_) + .verify(weight); + TensorMatcher({B, H, 1}) // + .with_dtype() + .with_device(device_) + .verify(weights_out); + TensorMatcher({-1, kRopeDim}) // + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({B}) // + .with_dtype(pos_dtype) + .with_device(device_) + .verify(positions); + + const auto batch_size = static_cast(B.unwrap()); + const auto num_heads = static_cast(H.unwrap()); + if (batch_size == 0) return; + + // The kernel computes row pointers as `base + work_id * kHeadDim`, so + // both inputs must be contiguous in (batch, head, elem) order. + const int64_t expected_batch_stride = static_cast(num_heads) * kHeadDim; + RuntimeCheck( + q_input.stride(0) == expected_batch_stride, + "q_input must be contiguous (B, H, kHeadDim); got stride[0]=", + q_input.stride(0)); + RuntimeCheck( + q_fp8.stride(0) == expected_batch_stride, + "q_fp8 must be contiguous (B, H, kHeadDim); got stride[0]=", + q_fp8.stride(0)); + + const auto params = FusedQIndexerRopeHadamardQuantParams{ + .q_input = q_input.data_ptr(), + .q_fp8 = q_fp8.data_ptr(), + .weight = weight.data_ptr(), + .weights_out = static_cast(weights_out.data_ptr()), + .weight_scale = static_cast(weight_scale), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .batch_size = batch_size, + .num_heads = num_heads, + }; + const auto total_works = batch_size * num_heads; + const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); + const auto k_int32 = kernel; + const auto k_int64 = kernel; + const auto k = pos_dtype.is_type() ? k_int32 : k_int64; + LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(k, params); + } +}; + +struct FusedQIndexerRopeHadamardFp4QuantParams { + const void* __restrict__ q_input; + void* __restrict__ q_fp4; + int32_t* __restrict__ q_sf; + const void* __restrict__ weight; + float* __restrict__ weights_out; + float weight_scale; + const float* __restrict__ freqs_cis; + const void* __restrict__ positions; + uint32_t batch_size; + uint32_t num_heads; +}; + +template +Q_KERNEL void +fused_q_indexer_rope_hadamard_fp4_quant(const __grid_constant__ FusedQIndexerRopeHadamardFp4QuantParams params) { + using namespace device; + + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kVecSize = 4; + constexpr uint32_t kRopeSize = kRopeDim / kVecSize; + static_assert(kHeadDim == kWarpThreads * kVecSize); + static_assert(kRopeDim == kWarpThreads * 2); + static_assert(kRopeSize <= kWarpThreads); + + using Storage = AlignedVector; + using Float4 = AlignedVector; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; + + const uint32_t total_works = params.batch_size * params.num_heads; + if (work_id >= total_works) return; + + const uint32_t batch_id = work_id / params.num_heads; + const auto input_ptr = static_cast(params.q_input) + work_id * kHeadDim; + const auto position = static_cast(static_cast(params.positions)[batch_id]); + const auto freqs_cis = params.freqs_cis + position * kRopeDim; + + PDLWaitPrimary(); + Float4 data, freq; + const auto weight_val = cast(static_cast(params.weight)[work_id]); + + { + Storage input_vec; + input_vec.load(input_ptr, lane_id); + if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + data[i] = cast(input_vec[i]); + } + } + + if (is_rope_lane) { + const auto x_real = data[0]; + const auto x_imag = data[1]; + const auto y_real = data[2]; + const auto y_imag = data[3]; + const auto fxr = freq[0]; + const auto fxi = freq[1]; + const auto fyr = freq[2]; + const auto fyi = freq[3]; + data[0] = x_real * fxr - x_imag * fxi; + data[1] = x_real * fxi + x_imag * fxr; + data[2] = y_real * fyr - y_imag * fyi; + data[3] = y_real * fyi + y_imag * fyr; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] = cast(cast(data[i])); + } + + PDLTriggerSecondary(); + + { + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a1; + data[1] = a0 - a1; + data[2] = a2 + a3; + data[3] = a2 - a3; + } + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a2; + data[1] = a1 + a3; + data[2] = a0 - a2; + data[3] = a1 - a3; + } +#pragma unroll + for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); + data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); + } + } + const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] *= kHadamardScale; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) + data[i] = cast(cast(data[i])); + } + + { + float local_max = math::abs(data[0]); +#pragma unroll + for (int i = 1; i < kVecSize; ++i) { + local_max = math::max(local_max, math::abs(data[i])); + } + local_max = warp::reduce_max<8>(local_max); + const auto scale_raw = fmaxf(1e-4f, local_max) / 6.0f; + const auto scale_ue8m0 = static_cast(cast_to_ue8m0(scale_raw)); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + const uint8_t packed0 = quant_fp4_e2m1(data[0] * inv_scale) | (quant_fp4_e2m1(data[1] * inv_scale) << 4); + const uint8_t packed1 = quant_fp4_e2m1(data[2] * inv_scale) | (quant_fp4_e2m1(data[3] * inv_scale) << 4); + const uint16_t packed = static_cast(packed0) | (static_cast(packed1) << 8); + auto out_row = static_cast(params.q_fp4) + work_id * (kHeadDim / 2); + reinterpret_cast(out_row)[lane_id] = packed; + if ((lane_id & 7) == 0) { + reinterpret_cast(params.q_sf + work_id)[lane_id >> 3] = scale_ue8m0; + } + params.weights_out[work_id] = weight_val * params.weight_scale; + } +} + +template +struct FusedQIndexerRopeHadamardFp4QuantKernel { + template + static constexpr auto kernel = fused_q_indexer_rope_hadamard_fp4_quant; + + static void forward( + const tvm::ffi::TensorView q_input, + const tvm::ffi::TensorView q_fp4, + const tvm::ffi::TensorView q_sf, + const tvm::ffi::TensorView weight, + const tvm::ffi::TensorView weights_out, + double weight_scale, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions) { + using namespace host; + constexpr int64_t kHeadDim = 128; + constexpr int64_t kRopeDim = 64; + constexpr int64_t kFp4Dim = kHeadDim / 2; + + auto B = SymbolicSize{"batch_size"}; + auto H = SymbolicSize{"num_heads"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, H, kHeadDim}) + .with_strides({-1, kHeadDim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_input); + TensorMatcher({B, H, kFp4Dim}) + .with_strides({-1, kFp4Dim, 1}) + .with_dtype() + .with_device(device_) + .verify(q_fp4); + TensorMatcher({B, H}).with_dtype().with_device(device_).verify(q_sf); + TensorMatcher({B, H}).with_dtype().with_device(device_).verify(weight); + TensorMatcher({B, H, 1}).with_dtype().with_device(device_).verify(weights_out); + TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); + auto pos_dtype = SymbolicDType{}; + TensorMatcher({B}).with_dtype(pos_dtype).with_device(device_).verify(positions); + + const auto batch_size = static_cast(B.unwrap()); + const auto num_heads = static_cast(H.unwrap()); + if (batch_size == 0) return; + + const int64_t expected_q_stride = static_cast(num_heads) * kHeadDim; + const int64_t expected_fp4_stride = static_cast(num_heads) * kFp4Dim; + RuntimeCheck(q_input.stride(0) == expected_q_stride, "q_input must be contiguous"); + RuntimeCheck(q_fp4.stride(0) == expected_fp4_stride, "q_fp4 must be contiguous"); + RuntimeCheck(q_sf.stride(0) == static_cast(num_heads) && q_sf.stride(1) == 1, "q_sf must be contiguous"); + + const auto params = FusedQIndexerRopeHadamardFp4QuantParams{ + .q_input = q_input.data_ptr(), + .q_fp4 = q_fp4.data_ptr(), + .q_sf = static_cast(q_sf.data_ptr()), + .weight = weight.data_ptr(), + .weights_out = static_cast(weights_out.data_ptr()), + .weight_scale = static_cast(weight_scale), + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .batch_size = batch_size, + .num_heads = num_heads, + }; + const auto total_works = batch_size * num_heads; + const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); + const auto k_int32 = kernel; + const auto k_int64 = kernel; + const auto k = pos_dtype.is_type() ? k_int32 : k_int64; + LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(k, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh new file mode 100644 index 0000000000..7d5f97824b --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh @@ -0,0 +1,219 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include + +namespace { + +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +struct MegaMoEPreDispatchParams { + const bf16_t* __restrict__ x; // [num_tokens, hidden] + const int32_t* __restrict__ topk_idx; // [num_tokens, top_k] + const float* __restrict__ topk_weights; // [num_tokens, top_k] + + fp8_e4m3_t* __restrict__ buf_x; // [padded_max, hidden] + int32_t* __restrict__ buf_x_sf; // contiguous int32 [P, G/4]; see layout comment + int64_t* __restrict__ buf_topk_idx; // [padded_max, top_k] + float* __restrict__ buf_topk_weights; // [padded_max, top_k] + + uint32_t num_tokens; + uint32_t padded_max; + uint32_t hidden; + uint32_t num_groups; // hidden / group_size + uint32_t top_k; +}; + +// kGroupSize must match sglang_per_token_group_quant_fp8_ue8m0(group_size=). +template +__global__ __launch_bounds__(1024, 2) void // + mega_moe_pre_dispatch_kernel(const MegaMoEPreDispatchParams __grid_constant__ params) { + using namespace device; + + constexpr uint32_t kVecElems = 8; // 8 bf16 = 16B load per thread + static_assert(kGroupSize % kVecElems == 0, "group_size must be a multiple of 8"); + constexpr uint32_t kThreadsPerGroup = kGroupSize / kVecElems; + using InputVec = AlignedVector; + using OutputVec = AlignedVector; + + const uint32_t bid = blockIdx.x; + const uint32_t tid = threadIdx.x; + + PDLWaitPrimary(); + if (bid < params.num_tokens) { + // ---- Quantize path: one CTA per valid token ---- + + const uint32_t token_id = bid; + const auto token_in = params.x + static_cast(token_id) * params.hidden; + const auto token_out = params.buf_x + static_cast(token_id) * params.hidden; + + InputVec in_vec; + in_vec.load(token_in, tid); + + float local_max = 0.0f; + float vals[kVecElems]; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + const auto [v0, v1] = cast(in_vec[i]); + vals[2 * i + 0] = v0; + vals[2 * i + 1] = v1; + local_max = fmaxf(local_max, fmaxf(fabsf(v0), fabsf(v1))); + } + + // Absmax across the kThreadsPerGroup threads that cover one group. + local_max = warp::reduce_max(local_max); + + const float absmax = fmaxf(local_max, 1e-10f); + const float raw_scale = absmax / math::FP8_E4M3_MAX; + const uint32_t ue8m0_exp = cast_to_ue8m0(raw_scale); + // 2^-ue8m0_exp as fp32 (equivalent to 1 / __uint_as_float(ue8m0 << 23)). + const float inv_scale = __uint_as_float((127u + 127u - ue8m0_exp) << 23); + + OutputVec out_vec; +#pragma unroll + for (uint32_t i = 0; i < kVecElems / 2; ++i) { + out_vec[i] = pack_fp8(vals[2 * i + 0] * inv_scale, vals[2 * i + 1] * inv_scale); + } + out_vec.store(token_out, tid); + + // One thread per group writes its UE8M0 byte into the contiguous + // row-major int32-packed layout: byte address = t*num_groups + g + // (see layout comment at the top of the file). + const uint32_t group_id = tid / kThreadsPerGroup; + const uint32_t within_group_id = tid % kThreadsPerGroup; + if (within_group_id == 0 && group_id < params.num_groups) { + const uint32_t byte_off = token_id * params.num_groups + group_id; + reinterpret_cast(params.buf_x_sf)[byte_off] = static_cast(ue8m0_exp); + } + + // Copy this token's topk row (no alignment assumptions; top_k is small). + if (tid < params.top_k) { + const uint32_t off = token_id * params.top_k + tid; + params.buf_topk_idx[off] = params.topk_idx[off]; + params.buf_topk_weights[off] = params.topk_weights[off]; + } + } else { + // ---- Pad path: trailing blocks fill [num_tokens, padded_max) with (-1, 0) ---- + const uint32_t copy_bid = bid - params.num_tokens; + const uint32_t pad_base = params.num_tokens * params.top_k; + const uint32_t slot = pad_base + copy_bid * blockDim.x + tid; + const uint32_t total_slots = params.padded_max * params.top_k; + + if (slot < total_slots) { + params.buf_topk_idx[slot] = -1; + params.buf_topk_weights[slot] = 0.0f; + } + } + PDLTriggerSecondary(); +} + +// ---- Host wrapper +// ------------------------------------------------------------------------------------------------------------------------ + +template +struct MegaMoEPreDispatchKernel { + static_assert(kGroupSize == 32 || kGroupSize == 64 || kGroupSize == 128, "unsupported group_size"); + static constexpr auto kernel = mega_moe_pre_dispatch_kernel(kGroupSize), kUsePDL>; + + static void + run(const tvm::ffi::TensorView x, + const tvm::ffi::TensorView topk_idx, + const tvm::ffi::TensorView topk_weights, + const tvm::ffi::TensorView buf_x, + const tvm::ffi::TensorView buf_x_sf, + const tvm::ffi::TensorView buf_topk_idx, + const tvm::ffi::TensorView buf_topk_weights) { + using namespace host; + + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto P = SymbolicSize{"padded_max"}; + auto H = SymbolicSize{"hidden"}; + auto K = SymbolicSize{"top_k"}; + auto G4 = SymbolicSize{"num_groups_div_4"}; + device.set_options(); + + TensorMatcher({M, H}) // input x + .with_dtype() + .with_device(device) + .verify(x); + TensorMatcher({M, K}) // topk_idx + .with_dtype() + .with_device(device) + .verify(topk_idx); + TensorMatcher({M, K}) // topk_weights + .with_dtype() + .with_device(device) + .verify(topk_weights); + TensorMatcher({P, H}) // buf.x + .with_dtype() + .with_device(device) + .verify(buf_x); + // buf.x_sf is the contiguous row-major int32 view from DeepGEMM's mega + // symm buffer (DeepGEMM/csrc/apis/mega.hpp): shape (P, G/4), strides + // (G/4, 1). No explicit strides required -> TensorMatcher enforces + // is_contiguous(). + TensorMatcher({P, G4}) // buf_x_sf + .with_dtype() + .with_device(device) + .verify(buf_x_sf); + TensorMatcher({P, K}) // buf.topk_idx + .with_dtype() + .with_device(device) + .verify(buf_topk_idx); + TensorMatcher({P, K}) // buf.topk_weights + .with_dtype() + .with_device(device) + .verify(buf_topk_weights); + + const auto num_tokens = static_cast(M.unwrap()); + const auto padded_max = static_cast(P.unwrap()); + const auto hidden = static_cast(H.unwrap()); + const auto top_k = static_cast(K.unwrap()); + const auto num_groups_div_4 = static_cast(G4.unwrap()); + + RuntimeCheck(num_tokens <= padded_max, "num_tokens must not exceed padded_max"); + RuntimeCheck(hidden % kGroupSize == 0, "hidden must be a multiple of group_size"); + const auto num_groups = hidden / static_cast(kGroupSize); + RuntimeCheck(num_groups == num_groups_div_4 * 4u, "num_groups must be a multiple of 4"); + RuntimeCheck(hidden % 8u == 0, "hidden must be a multiple of 8 (16B bf16 loads)"); + const auto num_threads = hidden / 8u; + RuntimeCheck(num_threads <= 1024, "hidden too large for single-block-per-row quant"); + RuntimeCheck(num_threads >= top_k, "top_k must fit into one quant CTA"); + + const auto pad_slots = (padded_max - num_tokens) * top_k; + const uint32_t num_pad_blocks = pad_slots == 0 ? 0u : ((pad_slots + num_threads - 1u) / num_threads); + const auto num_total_blocks = num_tokens + num_pad_blocks; + + const auto params = MegaMoEPreDispatchParams{ + .x = static_cast(x.data_ptr()), + .topk_idx = static_cast(topk_idx.data_ptr()), + .topk_weights = static_cast(topk_weights.data_ptr()), + .buf_x = static_cast(buf_x.data_ptr()), + .buf_x_sf = static_cast(buf_x_sf.data_ptr()), + .buf_topk_idx = static_cast(buf_topk_idx.data_ptr()), + .buf_topk_weights = static_cast(buf_topk_weights.data_ptr()), + .num_tokens = num_tokens, + .padded_max = padded_max, + .hidden = hidden, + .num_groups = num_groups, + .top_k = top_k, + }; + + if (num_total_blocks == 0) return; + LaunchKernel(num_total_blocks, num_threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh new file mode 100644 index 0000000000..38be975558 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh @@ -0,0 +1,119 @@ +#include +#include + +#include +#include + +#include +#include + +namespace { + +constexpr uint32_t kBlockSize = 1024; +constexpr uint32_t kSplitKV = 256; // const for both SM90 and SM100 + +struct MetadataParams { + /// NOTE: batch_size > 0 + uint32_t batch_size; + uint32_t num_sm; + const uint32_t* __restrict__ context_lens; + uint32_t* __restrict__ schedule_metadata; + bool use_smem = true; +}; + +__global__ __launch_bounds__(kBlockSize, 1) // + void smxx_paged_mqa_logits_metadata(const MetadataParams params) { + using namespace device; + extern __shared__ uint32_t s_length[]; + static constexpr auto kNumWarps = kBlockSize / kWarpThreads; + static_assert(kNumWarps == kWarpThreads); + + const auto tx = threadIdx.x; + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + __shared__ uint32_t s_warp_sum[kNumWarps]; + + uint32_t local_sum = 0; + for (uint32_t i = tx; i < params.batch_size; i += kBlockSize) { + const auto length = params.context_lens[i]; + local_sum += (length + kSplitKV - 1) / kSplitKV; + if (params.use_smem) s_length[i] = length; + } + + s_warp_sum[warp_id] = warp::reduce_sum(local_sum); + __syncthreads(); + + const auto global_sum = warp::reduce_sum(s_warp_sum[lane_id]); + if (lane_id != 0) return; + + const auto length_ptr = params.use_smem ? s_length : params.context_lens; + + const auto avg = global_sum / params.num_sm; + const auto ret = global_sum % params.num_sm; + uint32_t q = 0; + uint32_t num_work = (length_ptr[0] + kSplitKV - 1) / kSplitKV; + uint32_t sum_work = num_work; + for (auto i = warp_id; i <= params.num_sm; i += kNumWarps) { + const auto target = i * avg + min(i, ret); + while (sum_work <= target) { + if (++q >= params.batch_size) break; + num_work = (length_ptr[q] + kSplitKV - 1) / kSplitKV; + sum_work += num_work; + } + if (q >= params.batch_size) { + params.schedule_metadata[2 * i + 0] = params.batch_size; + params.schedule_metadata[2 * i + 1] = 0; + } else { + // sum > target && (sum - length) <= target + params.schedule_metadata[2 * i + 0] = q; + params.schedule_metadata[2 * i + 1] = target - (sum_work - num_work); + } + } +} + +template +void setup_kernel_smem_once(host::DebugInfo where = {}) { + [[maybe_unused]] + static const auto result = [] { + const auto fptr = std::bit_cast(f); + return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); + }(); + host::RuntimeDeviceCheck(result, where); +} + +struct IndexerMetadataKernel { + static constexpr auto kMaxBatchSizeInSmem = 16384 * 2; // 128 KB smeme + static void run(tvm::ffi::TensorView seq_lens, tvm::ffi::TensorView metadata) { + using namespace host; + auto N = SymbolicSize{"batch_size"}; + auto M = SymbolicSize{"num_sm"}; + auto device = SymbolicDevice{}; + device.set_options(); + TensorMatcher({N}) // + .with_dtype() + .with_device(device) + .verify(seq_lens); + TensorMatcher({M, 2}) // + .with_dtype() + .with_device(device) + .verify(metadata); + const auto batch_size = static_cast(N.unwrap()); + const auto num_sm = static_cast(M.unwrap()) - 1; + RuntimeCheck(num_sm <= 1024); + const auto use_smem = batch_size <= kMaxBatchSizeInSmem; + const auto params = MetadataParams{ + .batch_size = batch_size, + .num_sm = num_sm, + .context_lens = static_cast(seq_lens.data_ptr()), + .schedule_metadata = static_cast(metadata.data_ptr()), + .use_smem = use_smem, + }; + constexpr auto kernel = smxx_paged_mqa_logits_metadata; + setup_kernel_smem_once(); + const auto smem = use_smem ? (batch_size + 1) * sizeof(uint32_t) : 0; + LaunchKernel(1, kBlockSize, device.unwrap(), smem)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh new file mode 100644 index 0000000000..2239d3972d --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh @@ -0,0 +1,169 @@ +#include +#include + +#include +#include +#include + +#include + +#include + +namespace { + +using DType = bf16_t; +constexpr int64_t kRopeDim = 64; +constexpr uint32_t kBlockSize = 128; +constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + +struct FusedQKRopeParams { + void* __restrict__ q; + void* __restrict__ k; + const float* __restrict__ freqs_cis; + const void* __restrict__ positions; + int64_t q_stride_batch; + int64_t k_stride_batch; + int64_t q_stride_head; + int64_t k_stride_head; + uint32_t num_q_heads; + uint32_t num_k_heads; + uint32_t batch_size; +}; + +template +__global__ __launch_bounds__(kBlockSize, 16) // + void deepseek_rope_kernel(const __grid_constant__ FusedQKRopeParams param) { + using namespace device; + using DType2 = packed_t; + + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto global_warp_id = blockIdx.x * kNumWarps + warp_id; + + const auto& [ + q, k, freqs_cis, positions, // + q_stride_batch, k_stride_batch, q_stride_head, k_stride_head, // + num_q_heads, num_k_heads, batch_size + ] = param; + + const auto num_total_heads = num_q_heads + num_k_heads; + const auto head_id = global_warp_id % num_total_heads; + const auto batch_id = global_warp_id / num_total_heads; + if (batch_id >= batch_size) return; + + const auto position = static_cast(positions)[batch_id]; + const auto is_q = head_id < num_q_heads; + const auto local_head = is_q ? head_id : (head_id - num_q_heads); + const auto stride_batch = is_q ? q_stride_batch : k_stride_batch; + const auto stride_head = is_q ? q_stride_head : k_stride_head; + const auto base_ptr = is_q ? q : k; + const auto input = static_cast(pointer::offset(base_ptr, batch_id * stride_batch, local_head * stride_head)); + + const auto freq_ptr = reinterpret_cast(freqs_cis + position * kRopeDim); + const auto [f_real, f_imag] = freq_ptr[lane_id]; + PDLWaitPrimary(); + + const auto data = input[lane_id]; + const auto [x_real, x_imag] = cast(data); + fp32x2_t output; + if constexpr (kInverse) { + // (a + bi) * (c - di) = (ac + bd) + (bc - ad)i + output = { + x_real * f_real + x_imag * f_imag, + x_imag * f_real - x_real * f_imag, + }; + } else { + // (a + bi) * (c + di) = (ac - bd) + (ad + bc)i + output = { + x_real * f_real - x_imag * f_imag, + x_real * f_imag + x_imag * f_real, + }; + } + input[lane_id] = cast(output); + + PDLTriggerSecondary(); +} + +template +struct FusedQKRopeKernel { + // 4 kernel variants: {forward, inverse} x {int32, int64} + static constexpr auto kernel_fwd_i32 = deepseek_rope_kernel; + static constexpr auto kernel_fwd_i64 = deepseek_rope_kernel; + static constexpr auto kernel_inv_i32 = deepseek_rope_kernel; + static constexpr auto kernel_inv_i64 = deepseek_rope_kernel; + + static void forward( + const tvm::ffi::TensorView q, + const tvm::ffi::Optional k, + const tvm::ffi::TensorView freqs_cis, + const tvm::ffi::TensorView positions, + bool inverse) { + using namespace host; + + auto B = SymbolicSize{"batch_size"}; + auto Q = SymbolicSize{"num_q_heads"}; + auto K = SymbolicSize{"num_k_heads"}; + constexpr auto D = kRopeDim; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, Q, D}) // + .with_strides({-1, -1, 1}) + .with_dtype() + .with_device(device_) + .verify(q); + if (k.has_value()) { + TensorMatcher({B, K, D}) // + .with_strides({-1, -1, 1}) + .with_dtype() + .with_device(device_) + .verify(k.value()); + } else { + K.set_value(0); + } + TensorMatcher({-1, D}) // + .with_dtype() + .with_device(device_) + .verify(freqs_cis); + + auto pos_dtype = SymbolicDType{}; + TensorMatcher({B}) // + .with_dtype(pos_dtype) + .with_device(device_) + .verify(positions); + const bool pos_i32 = pos_dtype.is_type(); + + const auto batch_size = static_cast(B.unwrap()); + if (batch_size == 0) return; + + const auto num_q_heads = static_cast(Q.unwrap()); + const auto num_k_heads = static_cast(K.unwrap()); + const auto num_total_heads = num_q_heads + num_k_heads; + const auto total_warps = batch_size * num_total_heads; + const auto num_blocks = div_ceil(total_warps, kNumWarps); + + const auto elem_size = static_cast(sizeof(DType)); + const auto params = FusedQKRopeParams{ + .q = q.data_ptr(), + .k = k ? k.value().data_ptr() : nullptr, + .freqs_cis = static_cast(freqs_cis.data_ptr()), + .positions = positions.data_ptr(), + .q_stride_batch = q.stride(0) * elem_size, + .k_stride_batch = k ? k.value().stride(0) * elem_size : 0, + .q_stride_head = q.stride(1) * elem_size, + .k_stride_head = k ? k.value().stride(1) * elem_size : 0, + .num_q_heads = num_q_heads, + .num_k_heads = num_k_heads, + .batch_size = batch_size, + }; + + // dispatch: {inverse} x {pos_i32} + using KernelType = decltype(kernel_fwd_i32); + const KernelType kernel = + inverse ? (pos_i32 ? kernel_inv_i32 : kernel_inv_i64) : (pos_i32 ? kernel_fwd_i32 : kernel_fwd_i64); + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh new file mode 100644 index 0000000000..be0e759445 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh @@ -0,0 +1,540 @@ +#include +#include + +#include +#include +#include +#include +#include +#include + +#include + +#include +#include +#include + +namespace { + +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +struct SiluMulQuantVarlenParams { + const bf16_t* __restrict__ input; + fp8_e4m3_t* __restrict__ output; + float* __restrict__ output_scale; + const int32_t* __restrict__ masked_m; + float swiglu_limit; // only read when kApplySwigluLimit=true + int64_t hidden_dim; + uint32_t num_tokens; + uint32_t num_experts; +}; + +constexpr uint32_t kMaxExperts = 256; + +struct alignas(16) CTAWork { + uint32_t expert_id; + uint32_t expert_token_id; + bool valid; +}; + +SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { + static_assert(device::kWarpThreads == 32); +#pragma unroll + for (uint32_t offset = 1; offset < 32; offset *= 2) { + uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); + if (lane_id >= offset) val += n; + } + return val; +} + +template +SGL_DEVICE fp32x2_t silu_and_mul(DType2 gate, DType2 up, float limit) { + using namespace device; + // refer to as implementation. TL;DR: must clamp in bf16 + // https://github.com/deepseek-ai/DeepGEMM/blob/7f2a703ed51ac1f7af07f5e1453b2d3267d37d50/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh#L984-L997 + if constexpr (kApplySwigluLimit) { + static_assert(std::is_same_v); + gate = __hmin2(gate, {limit, limit}); + up = __hmax2(up, {-limit, -limit}); + up = __hmin2(up, {limit, limit}); + } + const auto [g0, g1] = cast(gate); + const auto [u0, u1] = cast(up); + const auto silu0 = g0 / (1.0f + __expf(-g0)); + const auto silu1 = g1 / (1.0f + __expf(-g1)); + const float val0 = silu0 * u0; + const float val1 = silu1 * u1; + if constexpr (kPrecise) { // I don't know if we should enable this? + return {val0, val1}; + } else { + return cast(cast(fp32x2_t{val0, val1})); + } +} + +[[maybe_unused]] +SGL_DEVICE CTAWork get_work(const SiluMulQuantVarlenParams& params) { + // Preconditions: + // 1. blockDim.x >= params.num_experts + // 2. params.num_experts <= kMaxExperts + using namespace device; + static_assert(kWarpThreads == 32); + + static __shared__ uint32_t s_warp_sum[32]; + static __shared__ CTAWork result; + + result.valid = false; + + const uint32_t tx = threadIdx.x; + const uint32_t lane_id = tx % kWarpThreads; + const uint32_t warp_id = tx / kWarpThreads; + + const uint32_t val = tx < params.num_experts ? params.masked_m[tx] : 0u; + + // Per-warp inclusive scan of masked_m. + const uint32_t warp_inclusive = warp_inclusive_sum(lane_id, val); + const uint32_t warp_exclusive = warp_inclusive - val; + + // Write each warp total. + if (lane_id == kWarpThreads - 1) s_warp_sum[warp_id] = warp_inclusive; + __syncthreads(); + const auto tmp_val = lane_id < warp_id ? s_warp_sum[lane_id] : 0u; + const auto prefix_exclusive = warp::reduce_sum(tmp_val) + warp_exclusive; + const auto bx = blockIdx.x; + if (prefix_exclusive <= bx && bx < prefix_exclusive + val) { + result = {tx, bx - prefix_exclusive, true}; + } + __syncthreads(); + return result; +} + +template +__global__ __launch_bounds__(1024, 2) void // maximize occupancy + silu_mul_quant_varlen_kernel(const SiluMulQuantVarlenParams __grid_constant__ params) { + using namespace device; + + constexpr uint32_t kGroupSize = 128u; + constexpr uint32_t kWorkThreads = 16u; + // each thread will handle 8 elements + using InputVec = AlignedVector; + using OutputVec = AlignedVector; + static_assert(8 * kWorkThreads == 128, "Invalid tiling"); + static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0"); + + const auto [expert_id, token_id, valid] = get_work(params); + + if (!valid) return; + + const auto work_id = threadIdx.x / kWorkThreads; + + const auto offset = expert_id * params.num_tokens + token_id; + const auto input = params.input + offset * params.hidden_dim * 2; + const auto output = params.output + offset * params.hidden_dim; + [[maybe_unused]] + const auto output_scale = [&] { + const auto num_groups = params.hidden_dim / kGroupSize; + if constexpr (kTransposed) { + const auto base = reinterpret_cast(params.output_scale); + // Physical layout is [E, G//4, N] int32. Each int32 packs 4 consecutive + // group scales for the same token, so the byte address is: + // expert_offset + (group/4)*N*4 + token*4 + group%4 + return base + expert_id * num_groups * params.num_tokens + (work_id / 4u) * (params.num_tokens * 4u) + + token_id * 4u + (work_id % 4u); + } else { + return params.output_scale + offset * num_groups + work_id; + } + }(); + + PDLWaitPrimary(); + + InputVec gate_vec, up_vec; + if constexpr (kSwizzle) { + // gran=8 interleaved: every 16-element chunk on the N axis is + // [gate[0..7], up[0..7]]. Each thread handles 8 consecutive output + // elements, so its gate chunk lives at vec index 2*threadIdx.x and its + // up chunk at 2*threadIdx.x+1. + gate_vec.load(input, threadIdx.x * 2); + up_vec.load(input, threadIdx.x * 2 + 1); + } else { + gate_vec.load(input, threadIdx.x); + up_vec.load(input, threadIdx.x + blockDim.x); + } + + float local_max = 0.0f; + float results[8]; + +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const auto [x, y] = silu_and_mul(gate_vec[i], up_vec[i], params.swiglu_limit); + results[2 * i + 0] = x; + results[2 * i + 1] = y; + local_max = fmaxf(local_max, fmaxf(fabsf(x), fabsf(y))); + } + + local_max = warp::reduce_max(local_max); + + const float absmax = fmaxf(local_max, 1e-10f); + float scale; + uint32_t ue8m0_exp; + + if constexpr (kScaleUE8M0) { + const float raw_scale = absmax / math::FP8_E4M3_MAX; + ue8m0_exp = cast_to_ue8m0(raw_scale); + scale = __uint_as_float(ue8m0_exp << 23); + } else { + scale = absmax / math::FP8_E4M3_MAX; + } + const auto inv_scale = 1.0f / scale; + + OutputVec out_vec; +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const float scaled_val0 = results[2 * i + 0] * inv_scale; + const float scaled_val1 = results[2 * i + 1] * inv_scale; + out_vec[i] = pack_fp8(scaled_val0, scaled_val1); + } + + PDLTriggerSecondary(); + + out_vec.store(output, threadIdx.x); + if constexpr (kTransposed) { + *output_scale = ue8m0_exp; + } else { + *output_scale = scale; + } +} + +struct SiluAndMulClampParams { + const void* __restrict__ input; + void* __restrict__ output; + float swiglu_limit; +}; + +template +__global__ __launch_bounds__(1024, 2) void // maximize occupancy + silu_mul_clamp_kernel(const SiluAndMulClampParams __grid_constant__ params) { + using namespace device; + static_assert(sizeof(DType) == 2, "only fp16/bf16 supported"); + using DType2 = packed_t; + constexpr auto kVecSize = 16 / sizeof(DType); + static_assert(kVecSize % 2 == 0 && kVecSize > 0); + using Vec = AlignedVector; + const auto bid = blockIdx.x; + const auto tile = tile::Memory::cta(); + const float limit = params.swiglu_limit; + + PDLWaitPrimary(); + const auto gate = tile.load(params.input, bid * 2 + 0); + const auto up = tile.load(params.input, bid * 2 + 1); + Vec out; + +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + out[i] = cast(silu_and_mul(cast(gate[i]), cast(up[i]), limit)); + } + + tile.store(params.output, out, bid); + PDLTriggerSecondary(); +} + +// ---- Host wrapper +// ------------------------------------------------------------------------------------------------------------------------ + +template +struct SiluAndMulMaskedPostQuantKernel { + static_assert(kGroupSize == 128); + static constexpr auto kernel_normal = + silu_mul_quant_varlen_kernel; + static constexpr auto kernel_transposed = + silu_mul_quant_varlen_kernel; + + static void + run(const tvm::ffi::TensorView input, + const tvm::ffi::TensorView output, + const tvm::ffi::TensorView output_scale, + const tvm::ffi::TensorView masked_m, + const uint32_t topk, + const bool transposed, + const double swiglu_limit) { + using namespace host; + + auto device = SymbolicDevice{}; + auto E = SymbolicSize{"num_experts"}; + auto T = SymbolicSize{"num_tokens_padded"}; + auto D = SymbolicSize{"hidden_dim x 2"}; + auto N = SymbolicSize{"hidden_dim"}; + auto G = SymbolicSize{"num_groups"}; + device.set_options(); + + TensorMatcher({E, T, D}) // input + .with_dtype() + .with_device(device) + .verify(input); + TensorMatcher({E, T, N}) // output + .with_dtype() + .with_device(device) + .verify(output); + if (!transposed) { + TensorMatcher({E, T, G}) // + .with_dtype() + .with_device(device) + .verify(output_scale); + } else { + RuntimeCheck(kScaleUE8M0, "transposed layout only supports scale_ue8m0=true"); + auto G_ = SymbolicSize{"G // 4"}; + TensorMatcher({E, G_, T}) // + .with_dtype() + .with_device(device) + .verify(output_scale); + G.set_value(G_.unwrap() * 4); + } + TensorMatcher({E}) // + .with_dtype() + .with_device(device) + .verify(masked_m); + + const auto num_experts = static_cast(E.unwrap()); + const auto num_tokens = static_cast(T.unwrap()); + const auto num_groups = static_cast(G.unwrap()); + const auto hidden_dim = N.unwrap(); + + RuntimeCheck(D.unwrap() == 2 * hidden_dim, "invalid dimension"); + RuntimeCheck(hidden_dim % kGroupSize == 0); + RuntimeCheck(num_experts <= kMaxExperts, "num_experts exceeds maximum (256)"); + RuntimeCheck(num_groups * kGroupSize == hidden_dim, "invalid num_groups"); + + const auto params = SiluMulQuantVarlenParams{ + .input = static_cast(input.data_ptr()), + .output = static_cast(output.data_ptr()), + .output_scale = static_cast(output_scale.data_ptr()), + .masked_m = static_cast(masked_m.data_ptr()), + .swiglu_limit = static_cast(swiglu_limit), + .hidden_dim = hidden_dim, + .num_tokens = num_tokens, + .num_experts = num_experts, + }; + + const auto num_threads = hidden_dim / 8; + RuntimeCheck(num_threads % device::kWarpThreads == 0); + RuntimeCheck(num_threads >= num_experts); + const auto kernel = transposed ? kernel_transposed : kernel_normal; + LaunchKernel(num_tokens * topk, num_threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +template +struct SiluAndMulClampKernel { + static constexpr auto kernel = silu_mul_clamp_kernel; + + static void run(const tvm::ffi::TensorView input, const tvm::ffi::TensorView output, const double swiglu_limit) { + using namespace host; + + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto D = SymbolicSize{"gate_up_dim"}; // 2 * out_dim + auto H = SymbolicSize{"out_dim"}; + device.set_options(); + + TensorMatcher({M, D}) // input (gate || up) + .with_dtype() + .with_device(device) + .verify(input); + TensorMatcher({M, H}) // output + .with_dtype() + .with_device(device) + .verify(output); + RuntimeCheck(D.unwrap() == 2 * H.unwrap(), "input last dim must be 2 * output last dim"); + + constexpr uint32_t kVecSize = 16 / sizeof(DType); + const auto out_dim = static_cast(H.unwrap()); + const auto num_tokens = static_cast(M.unwrap()); + RuntimeCheck(out_dim % kVecSize == 0, "out_dim must be divisible by vector size"); + const auto num_threads = out_dim / kVecSize; + RuntimeCheck(num_threads <= 1024, "out_dim too large for single-block-per-row launch"); + + const auto params = SiluAndMulClampParams{ + .input = input.data_ptr(), + .output = output.data_ptr(), + .swiglu_limit = static_cast(swiglu_limit), + }; + LaunchKernel(num_tokens, num_threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +struct SiluMulQuantContigParams { + const bf16_t* __restrict__ input; + fp8_e4m3_t* __restrict__ output; + float* __restrict__ output_scale; + float swiglu_limit; // only read when kApplySwigluLimit=true + int64_t hidden_dim; + uint32_t num_tokens; + uint32_t scale_row_stride_int32; // only used when kTransposed=true +}; + +template +__global__ __launch_bounds__(1024, 2) void // maximize occupancy + silu_mul_quant_contig_kernel(const SiluMulQuantContigParams __grid_constant__ params) { + using namespace device; + + constexpr uint32_t kGroupSize = 128u; + constexpr uint32_t kWorkThreads = 16u; + using InputVec = AlignedVector; + using OutputVec = AlignedVector; + static_assert(8 * kWorkThreads == 128, "Invalid tiling"); + static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0"); + + const auto token_id = blockIdx.x; + const auto work_id = threadIdx.x / kWorkThreads; + + const auto input = params.input + token_id * params.hidden_dim * 2; + const auto output = params.output + token_id * params.hidden_dim; + [[maybe_unused]] + const auto output_scale = [&] { + const auto num_groups = params.hidden_dim / kGroupSize; + if constexpr (kTransposed) { + // Physical layout is (G//4_pad, M_pad) int32; each int32 packs 4 + // consecutive UE8M0 exponents for the same token. Byte address: + // (work_id / 4) * M_pad * 4 + token * 4 + (work_id % 4). + const auto base = reinterpret_cast(params.output_scale); + return base + (work_id / 4u) * (params.scale_row_stride_int32 * 4u) + token_id * 4u + (work_id % 4u); + } else { + return params.output_scale + token_id * num_groups + work_id; + } + }(); + + PDLWaitPrimary(); + + InputVec gate_vec, up_vec; + if constexpr (kSwizzle) { + gate_vec.load(input, threadIdx.x * 2); + up_vec.load(input, threadIdx.x * 2 + 1); + } else { + gate_vec.load(input, threadIdx.x); + up_vec.load(input, threadIdx.x + blockDim.x); + } + + float local_max = 0.0f; + float results[8]; + +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const auto [x, y] = silu_and_mul(gate_vec[i], up_vec[i], params.swiglu_limit); + results[2 * i + 0] = x; + results[2 * i + 1] = y; + local_max = fmaxf(local_max, fmaxf(fabsf(x), fabsf(y))); + } + + local_max = warp::reduce_max(local_max); + + const float absmax = fmaxf(local_max, 1e-10f); + float scale; + uint32_t ue8m0_exp; + + if constexpr (kScaleUE8M0) { + const float raw_scale = absmax / math::FP8_E4M3_MAX; + ue8m0_exp = cast_to_ue8m0(raw_scale); + scale = __uint_as_float(ue8m0_exp << 23); + } else { + scale = absmax / math::FP8_E4M3_MAX; + } + const auto inv_scale = 1.0f / scale; + + OutputVec out_vec; +#pragma unroll + for (uint32_t i = 0; i < 4; ++i) { + const float scaled_val0 = results[2 * i + 0] * inv_scale; + const float scaled_val1 = results[2 * i + 1] * inv_scale; + out_vec[i] = pack_fp8(scaled_val0, scaled_val1); + } + + PDLTriggerSecondary(); + + out_vec.store(output, threadIdx.x); + if constexpr (kTransposed) { + *output_scale = ue8m0_exp; + } else { + *output_scale = scale; + } +} + +template +struct SiluAndMulContigPostQuantKernel { + static_assert(kGroupSize == 128); + static constexpr auto kernel_normal = + silu_mul_quant_contig_kernel; + static constexpr auto kernel_transposed = + silu_mul_quant_contig_kernel; + + static void + run(const tvm::ffi::TensorView input, + const tvm::ffi::TensorView output, + const tvm::ffi::TensorView output_scale, + const bool transposed, + const double swiglu_limit) { + using namespace host; + + auto device = SymbolicDevice{}; + auto M = SymbolicSize{"num_tokens"}; + auto D = SymbolicSize{"hidden_dim x 2"}; + auto N = SymbolicSize{"hidden_dim"}; + auto G = SymbolicSize{"num_groups"}; + device.set_options(); + + TensorMatcher({M, D}) // input (gate/up, natural or gran=8 interleaved on last dim) + .with_dtype() + .with_device(device) + .verify(input); + TensorMatcher({M, N}) // fp8 output + .with_dtype() + .with_device(device) + .verify(output); + + const auto hidden_dim = N.unwrap(); + RuntimeCheck(D.unwrap() == 2 * hidden_dim, "invalid dimension"); + RuntimeCheck(hidden_dim % kGroupSize == 0); + const auto num_groups = static_cast(hidden_dim / kGroupSize); + + uint32_t scale_row_stride_int32 = 0; + if (!transposed) { + G.set_value(num_groups); + TensorMatcher({M, G}) // (M, G) fp32 natural row-major + .with_dtype() + .with_device(device) + .verify(output_scale); + } else { + RuntimeCheck(kScaleUE8M0, "transposed layout only supports scale_ue8m0=true"); + RuntimeCheck(num_groups % 4 == 0, "transposed layout requires num_groups % 4 == 0"); + auto G_ = SymbolicSize{"G // 4"}; + G_.set_value(num_groups / 4); + auto M_pad = SymbolicSize{"M padded"}; + TensorMatcher({M, G_}) // `.transpose(-1,-2)[:M,:]` view of (G//4_pad, M_pad) int32 + .with_strides({int64_t{1}, M_pad}) // col-major transposed + .with_dtype() + .with_device(device) + .verify(output_scale); + scale_row_stride_int32 = static_cast(M_pad.unwrap()); + } + + const auto num_tokens = static_cast(M.unwrap()); + + const auto params = SiluMulQuantContigParams{ + .input = static_cast(input.data_ptr()), + .output = static_cast(output.data_ptr()), + .output_scale = static_cast(output_scale.data_ptr()), + .swiglu_limit = static_cast(swiglu_limit), + .hidden_dim = hidden_dim, + .num_tokens = num_tokens, + .scale_row_stride_int32 = scale_row_stride_int32, + }; + + const auto num_threads = hidden_dim / 8; + RuntimeCheck(num_threads % device::kWarpThreads == 0); + const auto kernel = transposed ? kernel_transposed : kernel_normal; + LaunchKernel(num_tokens, num_threads, device.unwrap()) // + .enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh new file mode 100644 index 0000000000..49f6f55963 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh @@ -0,0 +1,205 @@ +#include +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include + +#include +#include +#include + +namespace { + +using deepseek_v4::fp8::cast_to_ue8m0; +using deepseek_v4::fp8::inv_scale_ue8m0; +using deepseek_v4::fp8::pack_fp8; + +struct FusedStoreCacheParam { + const void* __restrict__ input; + void* __restrict__ cache; + const void* __restrict__ indices; + uint32_t num_tokens; +}; + +template +__global__ void fused_store_flashmla_cache(const __grid_constant__ FusedStoreCacheParam param) { + using namespace device; + + /// NOTE: 584 = 576 + 8 + constexpr int64_t kPageBytes = host::div_ceil(584 << kPageBits, 576) * 576; + + // each warp handles 64 elements, 8 warps, each block handles 1 row + const auto& [input, cache, indices, num_tokens] = param; + const uint32_t bid = blockIdx.x; + const uint32_t tid = threadIdx.x; + const uint32_t wid = tid / 32; + + PDLWaitPrimary(); + + // prefetch the index + const auto index = static_cast(indices)[bid]; + // always load the value from input (don't store if invalid) + using Float2 = packed_t; + const auto elems = static_cast(input)[tid + bid * 256]; + if (wid != 7) { + const auto [x, y] = cast(elems); + const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); + const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); + const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); + const auto result = pack_fp8(x * inv_scale, y * inv_scale); + const int32_t page = index >> kPageBits; + const int32_t offset = index & ((1 << kPageBits) - 1); + const auto page_ptr = pointer::offset(cache, page * kPageBytes); + const auto value_ptr = pointer::offset(page_ptr, offset * 576); + const auto scale_ptr = pointer::offset(page_ptr, 576 << kPageBits, offset * 8); + static_cast(value_ptr)[tid] = result; + static_cast(scale_ptr)[wid] = scale_ue8m0; + } else { + const auto result = cast(elems); + const int32_t page = index >> kPageBits; + const int32_t offset = index & ((1 << kPageBits) - 1); + const auto page_ptr = pointer::offset(cache, page * kPageBytes); + const auto value_ptr = pointer::offset(page_ptr, offset * 576, 448); + static_cast(value_ptr)[tid - 7 * 32] = result; + } + + PDLTriggerSecondary(); +} + +template +__global__ void fused_store_indexer_cache(const __grid_constant__ FusedStoreCacheParam param) { + using namespace device; + + /// NOTE: 132 = 128 + 4 + constexpr int64_t kPageBytes = 132 << kPageBits; + + // each warp handles 128 elements, 1 warp, each block handles multiple rows + const auto& [input, cache, indices, num_tokens] = param; + const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; + const auto global_wid = global_tid / 32; + const auto lane_id = threadIdx.x % 32; + + if (global_wid >= num_tokens) return; + + PDLWaitPrimary(); + + // prefetch the index + const auto index = static_cast(indices)[global_wid]; + // always load the value from input (don't store if invalid) + using Float2 = packed_t; + using InStorage = AlignedVector; + using OutStorage = AlignedVector; + const auto elems = static_cast(input)[global_tid]; + const auto [x0, x1] = cast(elems[0]); + const auto [y0, y1] = cast(elems[1]); + const auto local_max = fmaxf(fmaxf(fabs(x0), fabs(x1)), fmaxf(fabs(y0), fabs(y1))); + const auto abs_max = warp::reduce_max(local_max); + // use normal fp32 scale + const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; + const auto inv_scale = 1.0f / scale; + const int32_t page = index >> kPageBits; + const int32_t offset = index & ((1 << kPageBits) - 1); + const auto page_ptr = pointer::offset(cache, page * kPageBytes); + const auto value_ptr = pointer::offset(page_ptr, offset * 128); + const auto scale_ptr = pointer::offset(page_ptr, 128 << kPageBits, offset * 4); + OutStorage result; + result[0] = pack_fp8(x0 * inv_scale, x1 * inv_scale); + result[1] = pack_fp8(y0 * inv_scale, y1 * inv_scale); + static_cast(value_ptr)[lane_id] = result; + static_cast(scale_ptr)[0] = scale; + + PDLTriggerSecondary(); +} + +template +struct FusedStoreCacheFlashMLAKernel { + static constexpr int32_t kLogSize = std::countr_zero(kPageSize); + static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; + static constexpr auto kernel = fused_store_flashmla_cache; + + static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); + static_assert(1 << kLogSize == kPageSize); + + static void run(tvm::ffi::TensorView input, tvm::ffi::TensorView cache, tvm::ffi::TensorView indices) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({N, 512}) // input + .with_dtype() + .with_device(device_) + .verify(input); + TensorMatcher({-1, -1}) // cache + .with_strides({kPageBytes, 1}) + .with_dtype() + .with_device(device_) + .verify(cache); + TensorMatcher({N}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + const auto num_tokens = static_cast(N.unwrap()); + const auto params = FusedStoreCacheParam{ + .input = input.data_ptr(), + .cache = cache.data_ptr(), + .indices = indices.data_ptr(), + .num_tokens = num_tokens, + }; + const auto kBlockSize = 256; + const auto num_blocks = num_tokens; + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + } +}; + +template +struct FusedStoreCacheIndexerKernel { + static constexpr int32_t kLogSize = std::countr_zero(kPageSize); + static constexpr int64_t kPageBytes = 132 * kPageSize; + static constexpr auto kernel = fused_store_indexer_cache; + + static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); + static_assert(1 << kLogSize == kPageSize); + + static void run(tvm::ffi::TensorView input, tvm::ffi::TensorView cache, tvm::ffi::TensorView indices) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + TensorMatcher({N, 128}) // input + .with_dtype() + .with_device(device_) + .verify(input); + TensorMatcher({-1, -1}) // cache + .with_strides({kPageBytes, 1}) + .with_dtype() + .with_device(device_) + .verify(cache); + TensorMatcher({N}) // indices + .with_dtype() + .with_device(device_) + .verify(indices); + const auto num_tokens = static_cast(N.unwrap()); + const auto params = FusedStoreCacheParam{ + .input = input.data_ptr(), + .cache = cache.data_ptr(), + .indices = indices.data_ptr(), + .num_tokens = num_tokens, + }; + const auto kBlockSize = 128; + const auto num_blocks = div_ceil(num_tokens * 32, kBlockSize); + LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh new file mode 100644 index 0000000000..b1ccd24b20 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh @@ -0,0 +1,340 @@ +#include +#include + +#include + +#include +#include + +#include +#include + +namespace { + +#ifndef SGL_TOPK +#define SGL_TOPK 512 +#endif + +constexpr uint32_t kTopK = SGL_TOPK; +constexpr uint32_t kTopKBlockSize = SGL_TOPK; +constexpr uint32_t kSMEM = 16 * 1024 * sizeof(uint32_t); // 64KB (bytes) + +struct TopKParams { + const float* __restrict__ scores; + const int32_t* __restrict__ seq_lens; + const int32_t* __restrict__ page_table; + int32_t* __restrict__ page_indices; + int32_t* __restrict__ raw_indices; // optional: output raw abs position indices before page transform + const int64_t score_stride; + const int64_t page_table_stride; + uint32_t page_bits; +}; + +SGL_DEVICE uint8_t convert_to_uint8(float x) { + __half h = __float2half_rn(x); + uint16_t bits = __half_as_ushort(h); + uint16_t key = (bits & 0x8000) ? static_cast(~bits) : static_cast(bits | 0x8000); + return static_cast(key >> 8); +} + +SGL_DEVICE uint32_t convert_to_uint32(float x) { + uint32_t bits = __float_as_uint(x); + return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); +} + +SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) { + const uint32_t mask = (1u << page_bits) - 1u; + return (page_table[i >> page_bits] << page_bits) | (i & mask); +} + +[[maybe_unused]] +SGL_DEVICE void naive_transform( + const float* __restrict__, // unused + const int32_t* __restrict__ page_table, + int32_t* __restrict__ indices, + int32_t* __restrict__ raw_indices, // optional: output raw abs position indices + const uint32_t length, + const uint32_t page_bits) { + static_assert(kTopK <= kTopKBlockSize); + if (const auto tx = threadIdx.x; tx < length) { + indices[tx] = page_to_indices(page_table, tx, page_bits); + if (raw_indices != nullptr) { + raw_indices[tx] = tx; + } + } else if (kTopK == kTopKBlockSize || tx < kTopK) { + indices[tx] = -1; // fill invalid indices to -1 + if (raw_indices != nullptr) { + raw_indices[tx] = -1; + } + } +} + +[[maybe_unused]] +SGL_DEVICE void radix_topk(const float* __restrict__ input, int32_t* __restrict__ output, const uint32_t length) { + constexpr uint32_t RADIX = 256; + constexpr uint32_t BLOCK_SIZE = kTopKBlockSize; + constexpr uint32_t SMEM_INPUT_SIZE = kSMEM / (2 * sizeof(int32_t)); + + alignas(128) __shared__ uint32_t _s_histogram_buf[2][RADIX + 32]; + alignas(128) __shared__ uint32_t s_counter; + alignas(128) __shared__ uint32_t s_threshold_bin_id; + alignas(128) __shared__ uint32_t s_num_input[2]; + alignas(128) __shared__ int32_t s_last_remain; + + extern __shared__ uint32_t s_input_idx[][kSMEM / (2 * sizeof(int32_t))]; + + const uint32_t tx = threadIdx.x; + uint32_t remain_topk = kTopK; + auto& s_histogram = _s_histogram_buf[0]; + + const auto run_cumsum = [&] { +#pragma unroll 8 + for (int32_t i = 0; i < 8; ++i) { + static_assert(1 << 8 == RADIX); + if (tx < RADIX) { + const auto j = 1 << i; + const auto k = i & 1; + auto value = _s_histogram_buf[k][tx]; + if (tx + j < RADIX) { + value += _s_histogram_buf[k][tx + j]; + } + _s_histogram_buf[k ^ 1][tx] = value; + } + __syncthreads(); + } + }; + + // stage 1: 8bit coarse histogram + if (tx < RADIX + 1) s_histogram[tx] = 0; + __syncthreads(); + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const auto bin = convert_to_uint8(input[idx]); + ::atomicAdd(&s_histogram[bin], 1); + } + __syncthreads(); + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { + s_threshold_bin_id = tx; + s_num_input[0] = 0; + s_counter = 0; + } + __syncthreads(); + + const auto threshold_bin = s_threshold_bin_id; + remain_topk -= s_histogram[threshold_bin + 1]; + if (remain_topk == 0) { + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const uint32_t bin = convert_to_uint8(input[idx]); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } + } + __syncthreads(); + return; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const float raw_input = input[idx]; + const uint32_t bin = convert_to_uint8(raw_input); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } else if (bin == threshold_bin) { + const auto pos = ::atomicAdd(&s_num_input[0], 1); + if (pos < SMEM_INPUT_SIZE) { + [[likely]] s_input_idx[0][pos] = idx; + const auto bin = convert_to_uint32(raw_input); + const auto sub_bin = (bin >> 24) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + __syncthreads(); + } + + // stage 2: refine with 8bit radix passes +#pragma unroll 4 + for (int round = 0; round < 4; ++round) { + const auto r_idx = round % 2; + + // clip here to prevent overflow + const auto raw_num_input = s_num_input[r_idx]; + const auto num_input = raw_num_input < SMEM_INPUT_SIZE ? raw_num_input : SMEM_INPUT_SIZE; + + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { + s_threshold_bin_id = tx; + s_num_input[r_idx ^ 1] = 0; + s_last_remain = remain_topk - s_histogram[tx + 1]; + } + __syncthreads(); + + const auto threshold_bin = s_threshold_bin_id; + remain_topk -= s_histogram[threshold_bin + 1]; + + if (remain_topk == 0) { + for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(input[idx]) >> offset) & 0xFF; + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } + } + __syncthreads(); + break; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto raw_input = input[idx]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF; + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } else if (bin == threshold_bin) { + if (round == 3) { + const auto pos = ::atomicAdd(&s_last_remain, -1); + if (pos > 0) { + output[kTopK - pos] = idx; + } + } else { + const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1); + if (pos < SMEM_INPUT_SIZE) { + /// NOTE: (dark) fuse the histogram computation here + [[likely]] s_input_idx[r_idx ^ 1][pos] = idx; + const auto bin = convert_to_uint32(raw_input); + const auto sub_bin = (bin >> (offset - 8)) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + } + __syncthreads(); + } + } +} + +template +__global__ void topk_transform_kernel(const __grid_constant__ TopKParams params) { + const auto &[ + scores, seq_lens, page_table, page_indices, raw_indices, // pointers + score_stride, page_table_stride, page_bits // sizes + ] = params; + const uint32_t work_id = blockIdx.x; + + /// NOTE: dangerous prefetch seq_len before PDL wait + const uint32_t seq_len = seq_lens[work_id]; + const auto score_ptr = scores + work_id * score_stride; + const auto page_ptr = page_table + work_id * page_table_stride; + const auto indices_ptr = page_indices + work_id * kTopK; + const auto raw_indices_ptr = raw_indices != nullptr ? raw_indices + work_id * kTopK : nullptr; + + device::PDLWaitPrimary(); + + if (seq_len <= kTopK) { + naive_transform(score_ptr, page_ptr, indices_ptr, raw_indices_ptr, seq_len, page_bits); + } else { + __shared__ int32_t s_topk_indices[kTopK]; + radix_topk(score_ptr, s_topk_indices, seq_len); + static_assert(kTopK <= kTopKBlockSize); + const auto tx = threadIdx.x; + if (kTopK == kTopKBlockSize || tx < kTopK) { + indices_ptr[tx] = page_to_indices(page_ptr, s_topk_indices[tx], page_bits); + if (raw_indices_ptr != nullptr) { + raw_indices_ptr[tx] = s_topk_indices[tx]; + } + } + } + + device::PDLTriggerSecondary(); +} + +template +void setup_kernel_smem_once(host::DebugInfo where = {}) { + [[maybe_unused]] + static const auto result = [] { + const auto fptr = std::bit_cast(f); + return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); + }(); + host::RuntimeDeviceCheck(result, where); +} + +template +struct TopKKernel { + static constexpr auto kernel = topk_transform_kernel; + + static void transform( + const tvm::ffi::TensorView scores, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView page_table, + const tvm::ffi::TensorView page_indices, + const uint32_t page_size, + const tvm::ffi::Optional raw_indices) { + using namespace host; + auto B = SymbolicSize{"batch_size"}; + auto S = SymbolicSize{"score_stride"}; + auto P = SymbolicSize{"page_table_stride"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({B, -1}) // strided scores + .with_strides({S, 1}) + .with_dtype() + .with_device(device) + .verify(scores); + TensorMatcher({B}) // seq_lens, must be contiguous + .with_dtype() + .with_device(device) + .verify(seq_lens); + TensorMatcher({B, -1}) // strided page table + .with_strides({P, 1}) + .with_dtype() + .with_device(device) + .verify(page_table); + TensorMatcher({B, kTopK}) // output, must be contiguous + .with_dtype() + .with_device(device) + .verify(page_indices); + + int32_t* raw_indices_ptr = nullptr; + if (raw_indices.has_value()) { + TensorMatcher({B, kTopK}) // optional raw indices output, must be contiguous + .with_dtype() + .with_device(device) + .verify(raw_indices.value()); + raw_indices_ptr = static_cast(raw_indices.value().data_ptr()); + } + + RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2"); + const auto page_bits = static_cast(std::countr_zero(page_size)); + const auto batch_size = static_cast(B.unwrap()); + const auto params = TopKParams{ + .scores = static_cast(scores.data_ptr()), + .seq_lens = static_cast(seq_lens.data_ptr()), + .page_table = static_cast(page_table.data_ptr()), + .page_indices = static_cast(page_indices.data_ptr()), + .raw_indices = raw_indices_ptr, + .score_stride = S.unwrap(), + .page_table_stride = P.unwrap(), + .page_bits = page_bits, + }; + constexpr auto kSMEM_ = kSMEM + sizeof(int32_t); // align up a little + setup_kernel_smem_once(); + LaunchKernel(batch_size, kTopKBlockSize, device.unwrap(), kSMEM_).enable_pdl(kUsePDL)(kernel, params); + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh new file mode 100644 index 0000000000..8c4a526575 --- /dev/null +++ b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh @@ -0,0 +1,493 @@ +#include +#include + +#include +#include +#include +#include + +#include +#include +#include + +#include +#include +#include + +#include +#include +#include + +namespace { + +#ifndef SGL_TOPK +#define SGL_TOPK 512 +#endif + +inline constexpr uint32_t K = SGL_TOPK; + +template +void setup_kernel_smem_once(host::DebugInfo where = {}) { + [[maybe_unused]] + static const auto result = [] { + const auto fptr = std::bit_cast(f); + return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); + }(); + host::RuntimeDeviceCheck(result, where); +} + +namespace impl = device::top512; +using Large = impl::ClusterTopK; +using Medium = impl::StreamingTopK; +using Small = impl::RegisterTopK; + +using Metadata = Large::Metadata; +constexpr uint32_t kBlockSize = impl::kBlockSize; +constexpr uint32_t kNumClusters = 15; // based on hardware limits +constexpr uint32_t kClusterSize = Large::kClusterSize; +constexpr uint32_t kMax2PassLength = Small::kMax2PassLength; +constexpr uint32_t kMaxSupportedLength = Large::kMaxLength; + +/// Common metadata lives at metadata[0] (first row of the [batch_size+1, 4] tensor). +/// Per-item metadata starts at metadata[1..batch_size]. The plan kernel writes both. +struct alignas(16) GlobalMetadata { + uint32_t cluster_threshold; // decided per-batch in plan kernel + uint32_t num_cluster_items; // N = number of items routed to the cluster path + uint32_t reserved[2]; +}; +static_assert(sizeof(GlobalMetadata) == sizeof(Metadata), "layout: row 0 must occupy one Metadata-sized slot"); + +// optimize occupancy for prefill +#define SMALL_TOPK_KERNEL __global__ __launch_bounds__(kBlockSize, 2) +// cluster at y dim +#define LARGE_CLUSTER __cluster_dims__(1, kClusterSize, 1) +// stage-1 is persistent cluster, and shared memory usage is huge (can not 2) +#define LARGE_TOPK_STAGE_1 __global__ __launch_bounds__(kBlockSize, 1) LARGE_CLUSTER +// stage-2 is non-persistent non-cluster, with less shared memory and higher occupancy +#define LARGE_TOPK_STAGE_2 __global__ __launch_bounds__(kBlockSize, 2) +// fused into 1 stage when batch-size <= kNumPersistentClusters +#define FUSED_COMBINE_KERNEL __global__ __launch_bounds__(kBlockSize, 1) LARGE_CLUSTER +// plan runs once as a single block before the combine kernels +#define PLAN_KERNEL __global__ __launch_bounds__(kBlockSize, 1) + +struct TopKParams { + const uint32_t* __restrict__ seq_lens; + const float* __restrict__ scores; + const int32_t* __restrict__ page_table; + int32_t* __restrict__ page_indices; + int64_t score_stride; + int64_t page_table_stride; + uint8_t* __restrict__ workspace; // [batch, kWorkspaceBytes] -- internally allocated + /// Pointer to the full metadata tensor: metadata[0] is GlobalMetadata, metadata[1..] + /// are per-item entries (at most kNumClusters * rounds of them). + const Metadata* __restrict__ metadata = nullptr; + int64_t workspace_stride; // bytes per batch + uint32_t batch_size; + uint32_t page_bits; + + SGL_DEVICE const float* get_scores(const uint32_t batch_id) const { + return scores + batch_id * score_stride; + } + SGL_DEVICE impl::TransformParams get_transform(const uint32_t batch_id, int32_t* indices) const { + return { + .page_table = page_table + batch_id * page_table_stride, + .indices_in = indices, + .indices_out = page_indices + batch_id * K, + .page_bits = page_bits, + }; + } + SGL_DEVICE const GlobalMetadata& get_global_metadata() const { + return *reinterpret_cast(metadata); + } + SGL_DEVICE const Metadata& get_item_metadata(uint32_t work_id) const { + return metadata[1 + work_id]; // +1 to skip the GlobalMetadata row + } +}; + +SGL_DEVICE uint2 partition_work(uint32_t length, uint32_t rank) { + constexpr uint32_t kTMAAlign = 4; + const auto total_units = (length + kTMAAlign - 1) / kTMAAlign; + const auto base = total_units / kClusterSize; + const auto extra = total_units % kClusterSize; + const auto local_units = base + (rank < extra ? 1u : 0u); + const auto offset_units = rank * base + min(rank, extra); + const auto offset = offset_units * kTMAAlign; + const auto finish = min(offset + local_units * kTMAAlign, length); + return {offset, finish - offset}; +} + +/// Persistent scheduler. A single block: +/// 1. Decides a cluster_threshold from the real seq_lens distribution (or +/// uses the caller-supplied `static_cluster_threshold` when non-zero). +/// 2. Writes that threshold + N into metadata[0] (the GlobalMetadata row). +/// 3. Compacts items with seq_len > threshold into metadata[1..N+1), laid out +/// to match the persistent consumer's round-robin stride (kNumClusters). +/// Entries for clusters that get no work are zero-filled. +PLAN_KERNEL void topk_plan( + const uint32_t* __restrict__ seq_lens, + Metadata* __restrict__ metadata, + const uint32_t batch_size, + const uint32_t static_cluster_threshold) { + // Candidate thresholds, strictly increasing. Picked to give the auto-heuristic + // reasonable granularity without needing a full sort. Must all be >= kMax2PassLength. + + struct Pair { + uint32_t threshold; + uint32_t max_batch_size; + }; + /// NOTE: only tuned on B200 + constexpr Pair kCandidates[] = { + {32768, 30}, + {40960, 45}, + {49152, 45}, + {65536, 60}, + {98304, 60}, + {131072, 75}, + {196608, 90}, + {262144, 105}, + }; + constexpr uint32_t kNumCandidates = std::size(kCandidates); + constexpr uint32_t kMinBatchSize = kCandidates[0].max_batch_size; + static_assert(kCandidates[0].threshold == kMax2PassLength); + static_assert(kCandidates[kNumCandidates - 1].threshold == kMaxSupportedLength); + + __shared__ uint32_t s_count; // final N after compaction + __shared__ uint32_t s_counts[kNumCandidates]; + __shared__ uint32_t s_threshold; + + const auto tx = threadIdx.x; + if (tx == 0) s_count = 0; + if (tx < kNumCandidates) s_counts[tx] = 0; + __syncthreads(); + + // --- Phase 1: decide threshold ------------------------------------------ + if (static_cluster_threshold > 0) { + if (tx == 0) s_threshold = static_cluster_threshold; + } else if (batch_size <= kMinBatchSize) { + if (tx == 0) s_threshold = kMax2PassLength; // always prefer cluster + } else { + // Count items above each candidate threshold. Monotonically non-increasing in T. + for (uint32_t i = tx; i < batch_size; i += kBlockSize) { + const uint32_t sl = seq_lens[i]; + assert(sl <= kMaxSupportedLength); + uint32_t count = 0; +#pragma unroll + for (uint32_t j = 0; j < kNumCandidates; ++j) { + count += (sl > kCandidates[j].threshold ? 1 : 0); + } + if (count > 0) { + atomicAdd(&s_counts[count - 1], 1); + } + } + __syncthreads(); + if (tx == 0) { + uint32_t accum = 0; + uint32_t chosen = kMaxSupportedLength; +#pragma unroll + for (uint32_t i = 0; i < kNumCandidates; ++i) { + const auto j = kNumCandidates - 1 - i; + accum += s_counts[j]; + /// NOTE: `accum` increasing, while `max_batch_size` decreasing + if (accum > kCandidates[j].max_batch_size) break; + chosen = kCandidates[j].threshold; + } + s_threshold = chosen; + } + } + __syncthreads(); + // sanity check: below 2 pass threshold, must fits in small path + const auto cluster_threshold = max(s_threshold, kMax2PassLength); + + // --- Phase 2: compact items with seq_len > threshold into metadata[1..] - + // Per-item rows live at metadata[1 + pos]; metadata[0] is the GlobalMetadata row. + for (uint32_t i = tx; i < batch_size; i += kBlockSize) { + const uint32_t sl = seq_lens[i]; + if (sl > cluster_threshold) { + const auto pos = atomicAdd(&s_count, 1); + metadata[1 + pos] = {i, sl, false}; + } + } + __syncthreads(); + const auto N = s_count; + + // --- Phase 3: has_next + sentinels + GlobalMetadata --------------------- + for (uint32_t i = tx; i < N; i += kBlockSize) { + if (i + kNumClusters < N) metadata[1 + i].has_next = true; + } + // Zero-fill the first kNumClusters sentinel slots that got no valid entry. + if (tx < kNumClusters && tx >= N) metadata[1 + tx] = {0, 0, false}; + // Write global metadata (row 0). + if (tx == 0) { + auto* g = reinterpret_cast(metadata); + *g = { + .cluster_threshold = cluster_threshold, + .num_cluster_items = N, + .reserved = {0, 0}, + }; + } +} + +SMALL_TOPK_KERNEL void // short context +topk_short_transform(const __grid_constant__ TopKParams params) { + alignas(128) extern __shared__ uint8_t smem[]; + __shared__ int32_t s_topk_indices[K]; + const auto batch_id = blockIdx.x; + const auto seq_len = params.seq_lens[batch_id]; + const auto transform = params.get_transform(batch_id, s_topk_indices); + // trivial case + if (seq_len <= K) { + impl::trivial_transform(transform, seq_len, K); + } else { + Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem, /*use_pdl=*/true); + device::PDLTriggerSecondary(); + Small::transform(transform); + } +} + +LARGE_TOPK_STAGE_1 void // long context, middle to large batch size +topk_combine_preprocess(const __grid_constant__ TopKParams params) { + alignas(128) extern __shared__ uint8_t smem[]; + __shared__ int32_t s_topk_indices[K]; + uint32_t work_id = blockIdx.x; + uint32_t batch_id; + uint32_t seq_len; + bool has_next; + uint32_t length; + uint32_t offset; + const auto cluster_rank = blockIdx.y; + + const auto prefetch_metadata = [&] { + const auto metadata = params.get_item_metadata(work_id); + batch_id = metadata.batch_id; + seq_len = metadata.seq_len; + has_next = metadata.has_next; + work_id += kNumClusters; // advance to the next item for this cluster + }; + const auto launch_prologue = [&] { + const auto partition = partition_work(seq_len, cluster_rank); + offset = partition.x; + length = partition.y; + Large::stage1_prologue(params.get_scores(batch_id) + offset, length, smem); + }; + + device::PDLWaitPrimary(); + device::PDLTriggerSecondary(); + + prefetch_metadata(); + if (seq_len == 0) return; + Large::stage1_init(smem); + launch_prologue(); + while (true) { + const auto this_length = length; + const auto this_offset = offset; + const auto need_prefetch = has_next; + const auto transform = params.get_transform(batch_id, s_topk_indices); + const auto ws = params.workspace + batch_id * params.workspace_stride; + if (need_prefetch) prefetch_metadata(); + Large::stage1(s_topk_indices, this_length, smem, /*reuse=*/true); + if (need_prefetch) launch_prologue(); + Large::stage1_epilogue(transform, this_offset, ws, smem); + if (!need_prefetch) break; + } +} + +LARGE_TOPK_STAGE_2 void // long context, middle to large batch size +topk_combine_transform(const __grid_constant__ TopKParams params) { + alignas(128) extern __shared__ uint8_t smem[]; + __shared__ int32_t s_topk_indices[K]; + const auto batch_id = blockIdx.x; + const auto seq_len = params.seq_lens[batch_id]; + const auto cluster_threshold = params.get_global_metadata().cluster_threshold; + const auto transform = params.get_transform(batch_id, s_topk_indices); + if (seq_len <= K) { + impl::trivial_transform(transform, seq_len, K); + } else if (seq_len <= kMax2PassLength) { + if (seq_len <= Small::kMax1PassLength) { + Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem); + } else { + __syncwarp(); + Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem); + } + Small::transform(transform); + } else if (seq_len <= cluster_threshold) { + Medium::run(params.get_scores(batch_id), seq_len, s_topk_indices, smem); + Medium::transform(transform, smem); + } else { + const auto ws = params.workspace + batch_id * params.workspace_stride; + device::PDLWaitPrimary(); + Large::transform(transform, ws, smem); + } +} + +FUSED_COMBINE_KERNEL void // long context, small batch size +topk_fused_transform(const __grid_constant__ TopKParams params) { + alignas(128) extern __shared__ uint8_t smem[]; + __shared__ int32_t s_topk_indices[K]; + const auto batch_id = blockIdx.x; + const auto cluster_rank = blockIdx.y; + const auto seq_len = params.seq_lens[batch_id]; + const auto transform = params.get_transform(batch_id, s_topk_indices); + if (seq_len <= K) { + if (cluster_rank != 0) return; // only first rank work + impl::trivial_transform(transform, seq_len, K); + } else if (seq_len <= Small::kMax1PassLength) { + if (cluster_rank != 0) return; // only first rank work + Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem, /*use_pdl=*/true); + Small::transform(transform); + } else { + const auto [offset, length] = partition_work(seq_len, cluster_rank); + const auto ws = params.workspace + batch_id * params.workspace_stride; + Large::stage1_init(smem); + device::PDLWaitPrimary(); + Large::stage1_prologue(params.get_scores(batch_id) + offset, length, smem); + Large::stage1(s_topk_indices, length, smem); + Large::stage1_epilogue(transform, offset, ws, smem); + cooperative_groups::this_cluster().sync(); + if (cluster_rank != 0) return; // only first rank do the stage-2 + Large::transform(transform, ws, smem); + } +} + +struct CombinedTopKKernel { + static constexpr auto kStage1SMEM = sizeof(Large::Smem) + 128; + static constexpr auto kStage2SMEM = std::max(sizeof(Small::Smem), sizeof(Medium::Smem)) + 128; + + static void plan( // + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView metadata, + const uint32_t static_cluster_threshold) { + using namespace host; + auto B = SymbolicSize{"batch_size"}; + auto Bp1 = SymbolicSize{"batch_size_plus_1"}; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(seq_lens); + TensorMatcher({Bp1, 4}) // + .with_dtype() + .with_device(device_) + .verify(metadata); + + const auto batch_size = static_cast(B.unwrap()); + RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1); + if (batch_size <= kNumClusters) return; // metadata unused in fused path + + const auto device = device_.unwrap(); + constexpr auto kernel = topk_plan; + LaunchKernel(1, kBlockSize, device)( // + kernel, + static_cast(seq_lens.data_ptr()), + static_cast(metadata.data_ptr()), + batch_size, + static_cluster_threshold); + } + + static void transform( + const tvm::ffi::TensorView scores, + const tvm::ffi::TensorView seq_lens, + const tvm::ffi::TensorView page_table, + const tvm::ffi::TensorView page_indices, + const uint32_t page_size, + const tvm::ffi::TensorView workspace, + const tvm::ffi::TensorView metadata) { + using namespace host; + auto B = SymbolicSize{"batch_size"}; + auto Bp1 = SymbolicSize{"batch_size_plus_1"}; + auto L = SymbolicSize{"max_seq_len"}; + auto S = SymbolicSize{"score_stride"}; + auto P = SymbolicSize{"page_table_stride"}; + auto W = SymbolicSize{"workspace_stride"}; + constexpr auto D = Large::kWorkspaceInts; + auto device_ = SymbolicDevice{}; + device_.set_options(); + + TensorMatcher({B, L}) // + .with_strides({S, 1}) + .with_dtype() + .with_device(device_) + .verify(scores); + TensorMatcher({B}) // + .with_dtype() + .with_device(device_) + .verify(seq_lens); + TensorMatcher({B, -1}) // + .with_strides({P, 1}) + .with_dtype() + .with_device(device_) + .verify(page_table); + TensorMatcher({B, K}) // + .with_dtype() + .with_device(device_) + .verify(page_indices); + TensorMatcher({B, D}) // + .with_strides({W, 1}) + .with_dtype() + .with_device(device_) + .verify(workspace); + TensorMatcher({Bp1, 4}) // + .with_dtype() + .with_device(device_) + .verify(metadata); + + const auto page_bits = static_cast(std::countr_zero(page_size)); + const auto batch_size = static_cast(B.unwrap()); + const auto max_seq_len = static_cast(L.unwrap()); + const auto device = device_.unwrap(); + RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2"); + RuntimeCheck(S.unwrap() % 4 == 0, "score_stride must be a multiple of 4 (TMA 16-byte alignment)"); + RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1, "invalid metadata shape"); + + // NOTE: this should be fixed later + // RuntimeCheck(max_seq_len <= kMaxSupportedLength, max_seq_len, " exceeds the maximum supported length"); + + const auto params = TopKParams{ + .seq_lens = static_cast(seq_lens.data_ptr()), + .scores = static_cast(scores.data_ptr()), + .page_table = static_cast(page_table.data_ptr()), + .page_indices = static_cast(page_indices.data_ptr()), + .score_stride = S.unwrap(), + .page_table_stride = P.unwrap(), + .workspace = static_cast(workspace.data_ptr()), + .metadata = static_cast(metadata.data_ptr()), + .workspace_stride = W.unwrap() * static_cast(sizeof(int32_t)), + .batch_size = batch_size, + .page_bits = page_bits, + }; + + if (max_seq_len <= Small::kMax1PassLength) { + // All items fit in the short path -- no stage-1 needed + constexpr auto kernel = topk_short_transform; + setup_kernel_smem_once(); + LaunchKernel(batch_size, kBlockSize, device, kStage2SMEM) // + .enable_pdl(true)(kernel, params); + } else { + // Some items may be large -- launch stage-1 + main + if (batch_size <= kNumClusters) { + // can fuse into 1 stage + constexpr auto kernel = topk_fused_transform; + constexpr auto kSMEM = std::max(kStage1SMEM, kStage2SMEM); + setup_kernel_smem_once(); + LaunchKernel({batch_size, kClusterSize}, kBlockSize, device, kSMEM) + .enable_cluster({1, kClusterSize}) + .enable_pdl(true)(kernel, params); + } else { + // stage 1 + stage 2 + constexpr auto kernel_stage_1 = topk_combine_preprocess; + setup_kernel_smem_once(); + const auto num_clusters = std::min(batch_size, kNumClusters); + LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device, kStage1SMEM) + .enable_cluster({1, kClusterSize}) + .enable_pdl(true)(kernel_stage_1, params); + constexpr auto kernel_stage_2 = topk_combine_transform; + setup_kernel_smem_once(); + LaunchKernel(batch_size, kBlockSize, device, kStage2SMEM) // + .enable_pdl(true)(kernel_stage_2, params); + } + } + } +}; + +} // namespace diff --git a/lightllm/third_party/sglang_jit/dsv4/__init__.py b/lightllm/third_party/sglang_jit/dsv4/__init__.py new file mode 100644 index 0000000000..507b225167 --- /dev/null +++ b/lightllm/third_party/sglang_jit/dsv4/__init__.py @@ -0,0 +1,8 @@ +from .elementwise import fused_k_norm_rope_flashmla, fused_q_norm_rope +from .topk import topk_transform_512 + +__all__ = [ + "fused_k_norm_rope_flashmla", + "fused_q_norm_rope", + "topk_transform_512", +] diff --git a/lightllm/third_party/sglang_jit/dsv4/elementwise.py b/lightllm/third_party/sglang_jit/dsv4/elementwise.py new file mode 100644 index 0000000000..07011b0479 --- /dev/null +++ b/lightllm/third_party/sglang_jit/dsv4/elementwise.py @@ -0,0 +1,215 @@ +from typing import Optional, Tuple + +import torch + +from lightllm.third_party.sglang_jit.jit_utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from lightllm.third_party.sglang_jit.runtime_utils import is_hip + +from .utils import make_name + +_is_hip = is_hip() + + +@cache_once +def _jit_fused_rope_module(): + args = make_cpp_args(is_arch_support_pdl()) + return load_jit( + make_name("fused_rope"), + *args, + cuda_files=["deepseek_v4/rope.cuh"], + cuda_wrappers=[("forward", f"FusedQKRopeKernel<{args}>::forward")], + ) + + +@cache_once +def _jit_main_q_norm_rope_module( + dtype: torch.dtype, + head_dim: int, + rope_dim: int, +): + """Main MLA path Q kernel: rmsnorm-self + RoPE, warp per (token, head).""" + args = make_cpp_args(dtype, head_dim, rope_dim, is_arch_support_pdl()) + return load_jit( + make_name("main_q_norm_rope"), + *args, + cuda_files=["deepseek_v4/main_norm_rope.cuh"], + cuda_wrappers=[ + ("forward", f"FusedQNormRopeKernel<{args}>::forward"), + ], + ) + + +@cache_once +def _jit_main_k_norm_rope_flashmla_module( + dtype: torch.dtype, + head_dim: int, + rope_dim: int, + page_size: int, +): + """Main MLA path K kernel: rmsnorm + RoPE + write to FlashMLA paged cache.""" + args = make_cpp_args(dtype, head_dim, rope_dim, page_size, is_arch_support_pdl()) + return load_jit( + make_name("main_k_norm_rope_flashmla"), + *args, + cuda_files=["deepseek_v4/main_norm_rope.cuh"], + cuda_wrappers=[ + ("forward", f"FusedKNormRopeFlashMLAKernel<{args}>::forward"), + ], + ) + + +@cache_once +def _jit_main_q_indexer_rope_hadamard_quant_module(dtype: torch.dtype): + """C4 indexer Q kernel: RoPE + 128-pt Hadamard + fp8 act-quant""" + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + make_name("main_q_indexer_rope_hadamard_quant"), + *args, + cuda_files=["deepseek_v4/main_norm_rope.cuh"], + cuda_wrappers=[ + ("forward", f"FusedQIndexerRopeHadamardQuantKernel<{args}>::forward"), + ], + ) + + +@cache_once +def _jit_main_q_indexer_rope_hadamard_fp4_quant_module(dtype: torch.dtype): + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + make_name("main_q_indexer_rope_hadamard_fp4_quant"), + *args, + cuda_files=["deepseek_v4/main_norm_rope.cuh"], + cuda_wrappers=[ + ("forward", f"FusedQIndexerRopeHadamardFp4QuantKernel<{args}>::forward"), + ], + ) + + +def fused_rope_inplace( + q: torch.Tensor, + k: Optional[torch.Tensor], + freqs_cis: torch.Tensor, + positions: torch.Tensor, + inverse: bool = False, +) -> None: + """Apply rotary embeddings to both Q and K in a single fused CUDA kernel. + + Args: + q: [batch_size, num_q_heads, rope_dim] bfloat16 + k: [batch_size, num_k_heads, rope_dim] bfloat16 or None + freqs_cis: [max_seq_len, rope_dim // 2] complex64 (full table) + positions: [batch_size] int32 or int64, indices into freqs_cis + inverse: if True, apply inverse rotation (conjugate freqs) + """ + if _is_hip: + from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton + + apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse) + if k is not None: + apply_rotary_emb_triton(k, freqs_cis, positions=positions, inverse=inverse) + return + + freqs_real = torch.view_as_real(freqs_cis).flatten(-2).contiguous() + module = _jit_fused_rope_module() + module.forward(q, k, freqs_real, positions, inverse) + + +def fused_q_norm_rope( + q_input: torch.Tensor, + q_output: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +) -> None: + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + head_dim = q_input.shape[-1] + rope_dim = freqs_real.shape[-1] + module = _jit_main_q_norm_rope_module(q_input.dtype, head_dim, rope_dim) + module.forward(q_input, q_output, freqs_real, positions, eps) + + +def fused_q_indexer_rope_hadamard_quant( + q_input: torch.Tensor, + weight: torch.Tensor, + weight_scale: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + q_fp8 = torch.empty(q_input.shape, dtype=torch.float8_e4m3fn, device=q_input.device) + weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) + if _is_hip: + torch.ops.sgl_kernel.dsv4_fused_q_indexer_rope_hadamard_quant( + q_input, + q_fp8, + weight, + weights_out, + float(weight_scale), + freqs_real, + positions, + ) + else: + module = _jit_main_q_indexer_rope_hadamard_quant_module(q_input.dtype) + module.forward( + q_input, + q_fp8, + weight, + weights_out, + float(weight_scale), + freqs_real, + positions, + ) + return q_fp8, weights_out + + +def fused_q_indexer_rope_hadamard_fp4_quant( + q_input: torch.Tensor, + weight: torch.Tensor, + weight_scale: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor]: + if _is_hip: + raise RuntimeError("DeepSeek V4 FP4 indexer requires the CUDA fused Q path.") + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + q_fp4 = torch.empty( + (*q_input.shape[:-1], q_input.shape[-1] // 2), + dtype=torch.int8, + device=q_input.device, + ) + q_sf = torch.empty(q_input.shape[:-1], dtype=torch.int32, device=q_input.device) + weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) + module = _jit_main_q_indexer_rope_hadamard_fp4_quant_module(q_input.dtype) + module.forward( + q_input, + q_fp4, + q_sf, + weight, + weights_out, + float(weight_scale), + freqs_real, + positions, + ) + return (q_fp4, q_sf), weights_out + + +def fused_k_norm_rope_flashmla( + kv: torch.Tensor, + kv_weight: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, + out_loc: torch.Tensor, + kvcache: torch.Tensor, + page_size: int, +) -> None: + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + head_dim = kv.shape[-1] + rope_dim = freqs_real.shape[-1] + module = _jit_main_k_norm_rope_flashmla_module(kv.dtype, head_dim, rope_dim, page_size) + module.forward(kv, kv_weight, freqs_real, positions, out_loc, kvcache, eps) diff --git a/lightllm/third_party/sglang_jit/dsv4/topk.py b/lightllm/third_party/sglang_jit/dsv4/topk.py new file mode 100644 index 0000000000..1bfce7cef3 --- /dev/null +++ b/lightllm/third_party/sglang_jit/dsv4/topk.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +from typing import Optional + +import torch + +from lightllm.third_party.sglang_jit.jit_utils import ( + cache_once, + is_arch_support_pdl, + is_hip_runtime, + load_jit, + make_cpp_args, +) + +from .utils import make_name + + +@cache_once +def _jit_topk_v1_module(topk: int): + args = make_cpp_args(is_arch_support_pdl()) + assert topk in (512, 1024), "Only support topk=512 or 1024" + return load_jit( + make_name(f"topk_v1_{topk}"), + *args, + cuda_files=["deepseek_v4/topk_v1.cuh"], + cuda_wrappers=[("topk_transform", f"TopKKernel<{args}>::transform")], + extra_cuda_cflags=[f"-DSGL_TOPK={topk}"], + ) + + +@cache_once +def _jit_topk_v2_module(topk: int): + return load_jit( + make_name(f"topk_v2_{topk}"), + cuda_files=["deepseek_v4/topk_v2.cuh"], + cuda_wrappers=[ + ("topk_transform", "CombinedTopKKernel::transform"), + ("topk_plan", "CombinedTopKKernel::plan"), + ], + extra_cuda_cflags=[f"-DSGL_TOPK={topk}"], + ) + + +def topk_transform_512( + scores: torch.Tensor, + seq_lens: torch.Tensor, + page_tables: torch.Tensor, + out_page_indices: torch.Tensor, + page_size: int, + out_raw_indices: Optional[torch.Tensor] = None, +) -> None: + if is_hip_runtime(): + torch.ops.sgl_kernel.deepseek_v4_topk_transform_512( + scores, seq_lens, page_tables, out_page_indices, page_size, out_raw_indices + ) + else: + module = _jit_topk_v1_module(out_page_indices.shape[1]) + module.topk_transform(scores, seq_lens, page_tables, out_page_indices, page_size, out_raw_indices) + + +_WORKSPACE_INTS_PER_BATCH = 2 + 1024 * 2 +_PLAN_METADATA_INTS_PER_BATCH = 4 + + +def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor: + module = _jit_topk_v2_module(512) # does not matter + bs = seq_lens.shape[0] + metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH) + module.topk_plan(seq_lens, metadata, static_threshold) + return metadata + + +def topk_transform_512_v2( + scores: torch.Tensor, + seq_lens: torch.Tensor, + page_tables: torch.Tensor, + out_page_indices: torch.Tensor, + page_size: int, + metadata: torch.Tensor, +) -> None: + module = _jit_topk_v2_module(out_page_indices.shape[1]) + bs = scores.shape[0] + workspace = seq_lens.new_empty(bs, _WORKSPACE_INTS_PER_BATCH) + module.topk_transform( + scores, + seq_lens, + page_tables, + out_page_indices, + page_size, + workspace, + metadata, + ) diff --git a/lightllm/third_party/sglang_jit/dsv4/utils.py b/lightllm/third_party/sglang_jit/dsv4/utils.py new file mode 100644 index 0000000000..8085074f6c --- /dev/null +++ b/lightllm/third_party/sglang_jit/dsv4/utils.py @@ -0,0 +1,2 @@ +def make_name(name: str) -> str: + return f"dpsk_v4_{name}" diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh new file mode 100644 index 0000000000..c9da765f4a --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh @@ -0,0 +1,35 @@ +/// \file atomic.cuh +/// \brief Device-side atomic operations. + +#pragma once +#include + +namespace device::atomic { + +/** + * \brief Atomically computes the maximum of `*addr` and `value`, storing the + * result in `*addr`. + * \param addr Pointer to the value in global/shared memory to be updated. + * \param value The value to compare against. + * \return The old value at `*addr` before the update. + * \note On CUDA, this uses `atomicMax`/`atomicMin` on the reinterpreted + * integer representation. On ROCm, a CAS loop is used as a fallback. + */ +SGL_DEVICE float max(float* addr, float value) { +#ifndef USE_ROCM + float old; + old = (value >= 0) ? __int_as_float(atomicMax((int*)addr, __float_as_int(value))) + : __uint_as_float(atomicMin((unsigned int*)addr, __float_as_uint(value))); + return old; +#else + int* addr_as_i = (int*)addr; + int old = *addr_as_i, assumed; + do { + assumed = old; + old = atomicCAS(addr_as_i, assumed, __float_as_int(fmaxf(value, __int_as_float(assumed)))); + } while (assumed != old); + return __int_as_float(old); +#endif +} + +} // namespace device::atomic diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh new file mode 100644 index 0000000000..b47a4a27b2 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh @@ -0,0 +1,40 @@ +/// \file cta.cuh +/// \brief CTA (Cooperative Thread Array / thread-block) level primitives. + +#pragma once +#include +#include +#include + +namespace device::cta { + +/** + * \brief Compute the maximum of `value` across all threads in the CTA. + * + * Uses a two-level reduction: first within each warp via `warp::reduce_max`, + * then across warps using shared memory. The final result is stored in + * `smem[0]`. + * + * \tparam T Numeric type (must be supported by `warp::reduce_max`). + * \param value Per-thread input value. + * \param smem Shared memory buffer (must have at least `blockDim.x / 32` + * elements). + * \param min_value Identity element for max (default 0.0f). + * \note This function does NOT issue a trailing `__syncthreads()`. + * Callers must synchronize before reading `smem[0]`. + */ +template +SGL_DEVICE void reduce_max(T value, float* smem, float min_value = 0.0f) { + const uint32_t warp_id = threadIdx.x / kWarpThreads; + smem[warp_id] = warp::reduce_max(value); + __syncthreads(); + if (warp_id == 0) { + const auto tx = threadIdx.x; + const auto local_value = tx * kWarpThreads < blockDim.x ? smem[tx] : min_value; + const auto max_value = warp::reduce_max(local_value); + smem[0] = max_value; + } + // no extra sync; it is caller's responsibility to sync if needed +} + +} // namespace device::cta diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh new file mode 100644 index 0000000000..02b166d01c --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh @@ -0,0 +1,37 @@ +#pragma once + +#include + +#include + +#include +#include + +#include + +namespace device::compress { + +struct alignas(16) PrefillPlan { + uint32_t ragged_id; + uint32_t batch_id; + uint32_t position; + uint32_t window_len; // must be in `[0, compress_ratio * (1 + is_overlap))` + + bool is_valid(const uint32_t ratio, const bool is_overlap) const { + const uint32_t max_window_len = ratio * (1 + is_overlap); + return window_len < max_window_len; + } +}; + +} // namespace device::compress + +namespace host::compress { + +using device::compress::PrefillPlan; +using PrefillPlanTensorDtype = uint8_t; +inline constexpr int64_t kPrefillPlanDim = 16; + +static_assert(alignof(PrefillPlan) == sizeof(PrefillPlan)); +static_assert(sizeof(PrefillPlan) == kPrefillPlanDim * sizeof(PrefillPlanTensorDtype)); + +} // namespace host::compress diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh new file mode 100644 index 0000000000..3e87127c5f --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh @@ -0,0 +1,99 @@ +#pragma once + +#include +#include + +#include + +#include +#include + +#include + +namespace device::compress { + +/// \brief Per-batch decode plan. Layout: 16 bytes. +struct alignas(16) DecodePlan { + uint32_t seq_len; + int32_t write_loc; + int32_t read_page_0; + int32_t read_page_1; +}; + +/// \brief Per-token compress plan (used by c4/c128 prefill). Layout: 16 bytes. +struct alignas(16) CompressPlan { + uint32_t seq_len; + uint16_t ragged_id; + uint16_t buffer_len; + int32_t read_page_0; + /// \brief Stage 0 (CPU): batch_id (used to look up page table). + /// \brief Stage 1 (GPU): final state-pool write location. + int32_t read_page_1; + + static SGL_DEVICE __host__ CompressPlan invalid() { + return CompressPlan{-1u, 0, 0, -1, -1}; + } + + SGL_DEVICE __host__ bool is_invalid() const { + return seq_len == -1u; + } +}; + +/// \brief Per-token write plan (used by c4/c128 prefill). Layout: 8 bytes. +struct alignas(8) WritePlan { + /// \brief Stage 0 (CPU): packed `(batch_id << 16) | ragged_id`. + /// \brief Stage 1 (GPU): just `ragged_id`. + uint32_t ragged_id; + /// \brief Stage 0 (CPU): position + 1 (used to look up state slot). + /// \brief Stage 1 (GPU): final state-pool write location. + int32_t write_loc; + + static SGL_DEVICE __host__ WritePlan invalid() { + return WritePlan{-1u, -1}; + } + + SGL_DEVICE __host__ bool is_invalid() const { + return ragged_id == -1u; + } +}; + +} // namespace device::compress + +namespace host::compress { + +using device::compress::CompressPlan; +using device::compress::DecodePlan; +using device::compress::WritePlan; + +static_assert(alignof(DecodePlan) == sizeof(DecodePlan)); +static_assert(sizeof(DecodePlan) == 16); +static_assert(alignof(CompressPlan) == sizeof(CompressPlan)); +static_assert(sizeof(CompressPlan) == 16); +static_assert(alignof(WritePlan) == sizeof(WritePlan)); +static_assert(sizeof(WritePlan) == 8); + +inline auto verify_plan_d(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const DecodePlan* { + TensorMatcher({N, sizeof(DecodePlan)}) // + .with_dtype() + .with_device(device) + .verify(t); + return static_cast(t.data_ptr()); +} + +inline auto verify_plan_c(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const CompressPlan* { + TensorMatcher({N, sizeof(CompressPlan)}) // + .with_dtype() + .with_device(device) + .verify(t); + return static_cast(t.data_ptr()); +} + +inline auto verify_plan_w(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const WritePlan* { + TensorMatcher({N, sizeof(WritePlan)}) // + .with_dtype() + .with_device(device) + .verify(t); + return static_cast(t.data_ptr()); +} + +} // namespace host::compress diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh new file mode 100644 index 0000000000..53a62755b4 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh @@ -0,0 +1,112 @@ +#pragma once + +#include +#include +#include + +#include +#ifndef USE_ROCM +#include +#endif + +// Small helpers shared by the DeepSeek-V4 FP8/UE8M0 quantization kernels +// (silu_and_mul_masked_post_quant, store, mega_moe_pre_dispatch, ...). +// All functions are `SGL_DEVICE` (= `__forceinline__ __device__`) so +// including this header in multiple translation units is ODR-safe. + +namespace deepseek_v4::fp8 { + +// Round `x` to the nearest representable UE8M0 value. Returns the raw +// 8-bit biased exponent; the actual fp32 scale is `2^(exp - 127)` +// (i.e. `__uint_as_float(exp << 23)`). +SGL_DEVICE int32_t cast_to_ue8m0(float x) { + uint32_t u = __float_as_uint(x); + int32_t exp = int32_t((u >> 23) & 0xFF); + uint32_t mant = u & 0x7FFFFF; + return exp + (mant != 0); +} + +// 1 / 2^(exp - 127) as fp32. Equivalent to `1.0f / __uint_as_float(exp << 23)`. +SGL_DEVICE float inv_scale_ue8m0(int32_t exp) { + return __uint_as_float((127 + 127 - exp) << 23); +} + +// Clamp to [-FP8_E4M3_MAX, FP8_E4M3_MAX]. +// Uses platform-specific max from type.cuh (448 for E4M3FN, 224 for E4M3FNUZ). +SGL_DEVICE float fp8_e4m3_clip(float val) { + return fmaxf(fminf(val, kFP8E4M3Max), -kFP8E4M3Max); +} + +#ifndef USE_ROCM +// Pack two fp32 values into a single fp8x2_e4m3 with clamping. +SGL_DEVICE fp8x2_e4m3_t pack_fp8(float x, float y) { + return fp8x2_e4m3_t{fp32x2_t{fp8_e4m3_clip(x), fp8_e4m3_clip(y)}}; +} +#else +// Software float -> FP8 E4M3 conversion for ROCm/HIP. +// Supports both E4M3FN (MI350X, gfx950) and E4M3FNUZ (MI300X, gfx942). +SGL_DEVICE uint8_t cvt_float_to_fp8_e4m3(float val) { + val = fp8_e4m3_clip(val); + if (val == 0.0f) return 0; + + uint32_t f32 = __float_as_uint(val); + uint8_t sign = static_cast((f32 >> 31) << 7); + int32_t exp32 = static_cast((f32 >> 23) & 0xFF) - 127; + uint32_t mant23 = f32 & 0x7FFFFF; + +#if HIP_FP8_TYPE_FNUZ + // E4M3FNUZ: bias=8, max=240, no negative zero, NaN=0x80 + constexpr int32_t kBias = 8; + constexpr int32_t kMaxExp = 15; + constexpr int32_t kMinSubnormExp = -10; // min subnormal exponent + constexpr int32_t kMinNormExp = -7; // min normal exponent + constexpr uint8_t kSaturate = 0x7Fu; // max normal = 0_1111_111 = 240.0 +#else + // E4M3FN: bias=7, max=448, NaN=0x7F + constexpr int32_t kBias = 7; + constexpr int32_t kMaxExp = 15; + constexpr int32_t kMinSubnormExp = -9; + constexpr int32_t kMinNormExp = -6; + constexpr uint8_t kSaturate = 0x7Eu; // max normal = 0_1111_110 = 448.0 +#endif + + int32_t exp8; + uint8_t mant3; + + if (exp32 < kMinSubnormExp) { + return sign; + } else if (exp32 < kMinNormExp) { + // Subnormal range + int32_t shift = -(kBias - 1) - exp32; // 1..3 + uint32_t subnorm_mant = (0x800000 | mant23) >> (shift + 20); + uint32_t round_bit = ((0x800000 | mant23) >> (shift + 19)) & 1; + subnorm_mant += round_bit; + mant3 = static_cast(subnorm_mant & 0x07); + exp8 = 0; + if (subnorm_mant > 7) { + exp8 = 1; + mant3 = 0; + } + } else { + exp8 = exp32 + kBias; + mant3 = static_cast(mant23 >> 20); + uint32_t round_bit = (mant23 >> 19) & 1; + mant3 += round_bit; + if (mant3 > 7) { + mant3 = 0; + exp8++; + } + if (exp8 >= kMaxExp) return sign | kSaturate; + } + return sign | (static_cast(exp8) << 3) | mant3; +} + +// Pack two fp32 values into a single fp8x2_e4m3 (uint16_t on HIP). +SGL_DEVICE fp8x2_e4m3_t pack_fp8(float x, float y) { + uint8_t x8 = cvt_float_to_fp8_e4m3(x); + uint8_t y8 = cvt_float_to_fp8_e4m3(y); + return static_cast(x8) | (static_cast(y8) << 8); +} +#endif + +} // namespace deepseek_v4::fp8 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh new file mode 100644 index 0000000000..0a3acc4773 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh @@ -0,0 +1,96 @@ +#include +#include + +#include + +#include + +namespace device::hisparse { + +/// NOTE: We call nope+rope as a "value" here. +/// GPU Cache layout: +/// VALUE 0, VALUE 1, ..., VALUE 63, +/// SCALE 0, SCALE 1, ..., SCALE 63, +/// [Padding to align to 576 bytes] +/// CPU Cache follow a trivial linear layout without any padding. +inline constexpr int64_t kGPUPageSize = 64; +inline constexpr int64_t kGPUPageBits = 6; // log2(kGPUPageSize) +inline constexpr int64_t kValueBytes = 576; +inline constexpr int64_t kScaleBytes = 8; +/// NOTE: FlashMLA requires each page to be aligned to 576 bytes +inline constexpr int64_t kCPUItemBytes = kValueBytes + kScaleBytes; +inline constexpr int64_t kGPUPageBytes = host::div_ceil(kCPUItemBytes * kGPUPageSize, 576) * 576; +inline constexpr int64_t kGPUScaleOffset = kValueBytes * kGPUPageSize; + +struct PointerInfo { + int64_t* value_ptr; + int64_t* scale_ptr; +}; + +SGL_DEVICE PointerInfo get_pointer_gpu(void* cache, int32_t index) { + using namespace device; + static_assert(1 << kGPUPageBits == kGPUPageSize); + const int32_t page_num = index >> kGPUPageBits; + const int32_t page_offset = index & (kGPUPageSize - 1); + const auto page_ptr = pointer::offset(cache, page_num * kGPUPageBytes); + const auto value_ptr = pointer::offset(page_ptr, page_offset * kValueBytes); + const auto scale_ptr = pointer::offset(page_ptr, kGPUScaleOffset + page_offset * kScaleBytes); + return {static_cast(value_ptr), static_cast(scale_ptr)}; +} + +SGL_DEVICE PointerInfo get_pointer_cpu(void* cache, int32_t index) { + using namespace device; + const auto value_ptr = pointer::offset(cache, index * kCPUItemBytes); + const auto scale_ptr = pointer::offset(value_ptr, kValueBytes); + return {static_cast(value_ptr), static_cast(scale_ptr)}; +} + +enum class TransferDirection { + DeviceToDevice = 0, + DeviceToHost = 1, + HostToDevice = 2, +}; + +template +SGL_DEVICE void transfer_item(void* dst_cache, void* src_cache, const int32_t dst_index, const int32_t src_index) { + constexpr bool is_dst_device = (direction != TransferDirection::DeviceToHost); + constexpr bool is_src_device = (direction != TransferDirection::HostToDevice); + constexpr auto dst_fn = is_dst_device ? get_pointer_gpu : get_pointer_cpu; + constexpr auto src_fn = is_src_device ? get_pointer_gpu : get_pointer_cpu; + + const auto [dst_value_ptr, dst_scale_ptr] = dst_fn(dst_cache, dst_index); + const auto [src_value_ptr, src_scale_ptr] = src_fn(src_cache, src_index); + + int64_t local_items[2]; + const int64_t* tail_src_ptr; + int64_t* tail_dst_ptr; + + const int32_t lane_id = threadIdx.x % 32; + + for (int i = 0; i < 2; ++i) { + const auto j = lane_id + i * 32; + local_items[i] = src_value_ptr[j]; + } + + if (lane_id < 8) { // handle the tail element safely + const auto last_id = 64 + lane_id; + tail_src_ptr = src_value_ptr + last_id; + tail_dst_ptr = dst_value_ptr + last_id; + } else { // broadcast load/store is safe + tail_src_ptr = src_scale_ptr; + tail_dst_ptr = dst_scale_ptr; + } + + const auto tail_item = *tail_src_ptr; + + // store first 512 bytes of value + for (int i = 0; i < 2; ++i) { + const auto j = lane_id + i * 32; + dst_value_ptr[j] = local_items[i]; + } + + // store the tail element + *tail_dst_ptr = tail_item; +} + +} // namespace device::hisparse diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh new file mode 100644 index 0000000000..e58214c951 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh @@ -0,0 +1,257 @@ +#pragma once +#include +#include +#include + +#include "common.cuh" +#include "ptx.cuh" +#include +#include + +namespace device::top512 { + +template +struct ClusterTopK { + static constexpr uint32_t kClusterSize = 8; + static constexpr uint32_t kHistBits = 10; + static constexpr uint32_t kHistBins = 1 << kHistBits; + static constexpr uint32_t kRadixBins = 256; + static constexpr uint32_t kElemPerStage = 8; + static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize; + static constexpr uint32_t kNumStages = 4; + static constexpr uint32_t kMaxLength = kClusterSize * kNumStages * kSizePerStage; + static constexpr uint32_t kStoreLane = kBlockSize - 1; + static constexpr uint32_t kAboveBits = 11; + + // --------------------------------------------------------------------------- + // Shared memory layouts + // --------------------------------------------------------------------------- + + struct Smem { + uint64_t barrier[kNumStages]; + uint32_t local_above_equal[kClusterSize]; + uint32_t prefix_above_equal; + alignas(128) uint32_t counter_gt; + alignas(128) uint32_t counter_eq; + alignas(128) MatchBin match; + alignas(128) uint32_t warp_sum[kNumWarps]; + uint32_t histogram[kHistBins]; + alignas(128) float score_buffer[kNumStages][kSizePerStage]; + Tie tie_buffer[kMaxTies]; + }; + + struct alignas(16) Metadata { + uint32_t batch_id; + uint32_t seq_len; + bool has_next; + }; + + struct WorkSpace { + uint2 metadata; // {num_above, num_ties} + Tie ties[kMaxTies]; + }; + + static constexpr uint32_t kWorkspaceInts = sizeof(WorkSpace) / sizeof(uint32_t); + + // --------------------------------------------------------------------------- + // Stage 1: histogram + cluster reduce + find threshold + scatter + // --------------------------------------------------------------------------- + + SGL_DEVICE static void stage1_init(void* _smem) { + const auto tx = threadIdx.x; + __builtin_assume(tx < kBlockSize); + const auto smem = static_cast(_smem); + if (tx < kHistBins) smem->histogram[tx] = 0; + if (tx < kNumStages) ptx::mbarrier_init(&smem->barrier[tx], 1); + __syncthreads(); + } + + SGL_DEVICE static void stage1_prologue(const float* scores, uint32_t length, void* _smem) { + if (threadIdx.x == 0) { + const auto smem = static_cast(_smem); + const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage; + const auto length_aligned = (length + 3u) & ~3u; // align to 4 for TMA +#pragma unroll + for (uint32_t stage = 0; stage < kNumStages; stage++) { + if (stage >= num_stages) break; + const auto offset = stage * kSizePerStage; + const auto size = min(kSizePerStage, length_aligned - offset); + const auto size_bytes = size * sizeof(float); + const auto bar = &smem->barrier[stage]; + ptx::tma_load(smem->score_buffer[stage], scores + offset, size_bytes, bar); + ptx::mbarrier_arrive_expect_tx(bar, size_bytes); + } + } + } + + SGL_DEVICE static void stage1(int32_t* indices, uint32_t length, void* _smem, bool reuse = false) { + const auto smem = static_cast(_smem); + const auto tx = threadIdx.x; + __builtin_assume(tx < kBlockSize); + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + // Initialize shared memory histogram, counters, and barriers +#pragma unroll + for (uint32_t stage = 0; stage < kNumStages; stage++) { + const auto offset = stage * kSizePerStage; + if (offset >= length) break; + const auto size = min(kSizePerStage, length - offset); + if (lane_id == 0) ptx::mbarrier_wait(&smem->barrier[stage], 0); + __syncwarp(); +#pragma unroll + for (uint32_t i = 0; i < kElemPerStage; ++i) { + const auto idx = tx + i * kBlockSize; + if (idx >= size) break; + const auto score = smem->score_buffer[stage][idx]; + const auto bin = extract_coarse_bin(score); + atomicAdd(&smem->histogram[bin], 1); + } + } + + static_assert(kHistBins <= kBlockSize); + + // 2-shot all-reduce + { + auto cluster = cooperative_groups::this_cluster(); + cluster.sync(); + const auto cluster_rank = blockIdx.y; + const auto kLocalSize = kHistBins / kClusterSize; + const auto offset = kLocalSize * cluster_rank; + + const auto src_tx = tx / kClusterSize; + const auto src_rank = tx % kClusterSize; + + if (tx < kHistBins) { + const auto addr = &smem->histogram[offset + src_tx]; + const auto src_addr = cluster.map_shared_rank(addr, src_rank); + *src_addr = warp::reduce_sum(*src_addr); + } + cluster.sync(); + } + + // now each block holds the whole histogram, find the threshold bin + { + const auto value = tx < kHistBins ? smem->histogram[tx] : 0; + const auto warp_inc = warp_inclusive_sum(lane_id, value); + if (lane_id == kWarpThreads - 1) { + smem->warp_sum[warp_id] = warp_inc; + } + + __syncthreads(); + const auto tmp = smem->warp_sum[lane_id]; + // total_length = sum of all bins in the globally-reduced histogram + // (problem.length is block-local; after cluster reduction we need the global total) + const auto total_length = warp::reduce_sum(tmp); + uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); + prefix_sum += warp_inc; + const auto above = total_length - prefix_sum; + if (tx < kHistBins && above < K && above + value >= K) { + smem->counter_gt = smem->counter_eq = 0; + smem->match = { + .bin = tx, + .above_count = above, + .equal_count = value, + }; + } + __syncthreads(); + } + + const auto [thr_bin, num_above, num_equal] = smem->match; + + // write above and equal results to global memory +#pragma unroll + for (uint32_t stage = 0; stage < kNumStages; stage++) { + const auto offset = stage * kSizePerStage; + if (offset >= length) break; +#pragma unroll + for (uint32_t i = 0; i < kElemPerStage; ++i) { + const auto buf_idx = tx + i * kBlockSize; + const auto global_idx = offset + buf_idx; + if (global_idx >= length) break; + const auto score = smem->score_buffer[stage][buf_idx]; + const auto bin = extract_coarse_bin(score); + if (bin > thr_bin) { + indices[atomicAdd(&smem->counter_gt, 1)] = global_idx; + } else if (bin == thr_bin) { + const auto pos = atomicAdd(&smem->counter_eq, 1); + if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score}; + } + } + } + if (reuse) { + const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage; + if (tx < kHistBins) smem->histogram[tx] = 0; + if (tx < num_stages) ptx::mbarrier_arrive(&smem->barrier[tx]); + } + __syncthreads(); + } + + // --------------------------------------------------------------------------- + // Stage 1 epilogue: cross-block prefix sum + page translate + tie store + // --------------------------------------------------------------------------- + + SGL_DEVICE static void stage1_epilogue(const TransformParams params, const uint32_t offset, void* _ws, void* _smem) { + auto cluster = cooperative_groups::this_cluster(); + const auto smem = static_cast(_smem); + const auto tx = threadIdx.x; + const auto local_above = smem->counter_gt; + const auto local_equal = smem->counter_eq; + const auto cluster_rank = blockIdx.y; + + constexpr uint32_t kAboveMask = (1 << kAboveBits) - 1; + static_assert(kAboveMask >= K); + + // Pack local counts -- NO alignment rounding (contiguous layout) + static_assert(kMaxTies <= kBlockSize); + const auto idx_above = tx < local_above ? params.indices_in[tx] : 0; + const auto tie_value = tx < local_equal ? smem->tie_buffer[tx] : Tie{0, 0.0f}; + + // push to remote shared memory, can reduce latency of reading remote + if (tx < kClusterSize) { + const auto value = (local_equal << kAboveBits) | local_above; + const auto dst_addr = cluster.map_shared_rank(smem->local_above_equal, tx); + dst_addr[cluster_rank] = value; + } + // after this last sync, only read local shared memory + // so that it is safe when peer rank has already exited the kernel + cluster.sync(); + if (tx < kClusterSize) { + const auto value = tx < cluster_rank ? smem->local_above_equal[tx] : 0; + const auto kActiveMask = (1u << kClusterSize) - 1; + smem->prefix_above_equal = warp::reduce_sum(value, kActiveMask); + } + __syncthreads(); + + const auto prefix_packed = smem->prefix_above_equal; + const auto prefix_above = prefix_packed & kAboveMask; + const auto prefix_equal = prefix_packed >> kAboveBits; + + // Page-translate above elements + if (tx < local_above) { + params.write(tx + prefix_above, idx_above + offset); + } + // Contiguous tie store via regular global writes (no TMA, no gaps) + const auto ws = static_cast(_ws); + if (tx < local_equal && tx + prefix_equal < kMaxTies) { + ws->ties[tx + prefix_equal] = {tie_value.idx + offset, tie_value.score}; + } + // Block 0 writes global metadata {num_above, num_ties} + if (cluster_rank == kClusterSize - 1 && tx == 0) { + const auto sum_above = prefix_above + local_above; + const auto sum_equal = prefix_equal + local_equal; + ws->metadata = make_uint2(sum_above, sum_equal); + } + } + + SGL_DEVICE static void transform(const TransformParams params, const void* _ws, void* _smem) { + const auto ws = static_cast(_ws); + const auto meta = &ws->metadata; + const auto [num_above, num_equal] = *meta; + if (num_above >= K || num_equal == 0) return; + const auto clamped_ties = min(num_equal, kMaxTies); + tie_handle_transform(ws->ties, clamped_ties, num_above, K, params, _smem); + } +}; + +} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh new file mode 100644 index 0000000000..d553032d79 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh @@ -0,0 +1,176 @@ +#pragma once +#include +#include +#include +#include + +#include + +namespace device::top512 { + +inline constexpr uint32_t kMaxTopK = 1024; +inline constexpr uint32_t kBlockSize = 1024; +inline constexpr uint32_t kNumWarps = kBlockSize / kWarpThreads; +inline constexpr uint32_t kMaxTies = 1024; // == kBlockSize: 1 element per thread in stage2 +static constexpr uint32_t kRadixBins = 256; +static_assert(kMaxTopK <= kBlockSize && kMaxTies <= kBlockSize); + +// always use float4 to load from global memory +using Vec4 = AlignedVector; + +SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) { + const uint32_t mask = (1u << page_bits) - 1u; + return (page_table[i >> page_bits] << page_bits) | (i & mask); +} + +struct TransformParams { + const int32_t* __restrict__ page_table; + const int32_t* __restrict__ indices_in; + int32_t* __restrict__ indices_out; + uint32_t page_bits; + + SGL_DEVICE void transform(const uint32_t idx) const { + indices_out[idx] = page_to_indices(page_table, indices_in[idx], page_bits); + } + SGL_DEVICE void write(const uint32_t dst, const uint32_t src) const { + indices_out[dst] = page_to_indices(page_table, src, page_bits); + } +}; + +struct alignas(16) MatchBin { + uint32_t bin; + uint32_t above_count; + uint32_t equal_count; +}; + +struct alignas(8) Tie { + uint32_t idx; + float score; +}; + +struct TieHandleSmem { + alignas(128) uint32_t counter; // output position counter + alignas(128) MatchBin match; + uint32_t histogram[kRadixBins]; // 256-bin radix histogram + uint32_t warp_sum[kNumWarps]; // for 2-pass prefix sum +}; + +template +SGL_DEVICE uint32_t extract_coarse_bin(float x) { + static_assert(0 < kBits && kBits < 15); + const auto hx = cast(x); + const uint16_t bits = *reinterpret_cast(&hx); + const uint16_t key = (bits & 0x8000) ? ~bits : bits | 0x8000; + return key >> (16 - kBits); +} + +SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { + static_assert(kWarpThreads == 32); +#pragma unroll + for (uint32_t offset = 1; offset < 32; offset *= 2) { + uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); + if (lane_id >= offset) val += n; + } + return val; +} + +/// Order-preserving float32 -> uint32 for radix select +SGL_DEVICE uint32_t extract_exact_bin(float x) { + uint32_t bits = __float_as_uint(x); + return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); +} + +SGL_DEVICE void trivial_transform(const TransformParams& params, uint32_t length, uint32_t K) { + if (const auto tx = threadIdx.x; tx < length) { + params.write(tx, tx); + } else if (tx < K) { + params.indices_out[tx] = -1; + } +} + +SGL_DEVICE void tie_handle_transform( + const Tie* __restrict__ ties, // + const uint32_t num_ties, + const uint32_t num_above, + const uint32_t K, + const TransformParams params, + void* _smem) { + auto* smem = static_cast(_smem); + const auto tx = threadIdx.x; + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + // Each thread loads one element (or becomes inactive) + const bool has_elem = tx < num_ties; + const auto tie = has_elem ? ties[tx] : Tie{0, 0.0f}; + const uint32_t key = extract_exact_bin(tie.score); + const uint32_t idx = tie.idx; + bool active = has_elem; + uint32_t topk_remain = K - num_above; + uint32_t write_pos = K; + + smem->counter = 0; + __syncthreads(); + + // Number of warps covering the 256-bin histogram (256/32 = 8) + constexpr uint32_t kRadixWarps = kRadixBins / kWarpThreads; + +#pragma unroll + for (int round = 0; round < 4; round++) { + const uint32_t shift = 24 - round * 8; + const uint32_t bin = (key >> shift) & 0xFFu; + + // 1. Build histogram + if (tx < kRadixBins) smem->histogram[tx] = 0; + __syncthreads(); + if (active) atomicAdd(&smem->histogram[bin], 1); + __syncthreads(); + + // 2. v2-style 2-pass prefix sum on 256 bins + // Only first 256 threads (8 warps) carry histogram bins. + // Other threads get hist_val=0 and harmless prefix results. + uint32_t hist_val = 0; + uint32_t warp_inc = 0; + if (tx < kRadixBins) { + hist_val = smem->histogram[tx]; + warp_inc = warp_inclusive_sum(lane_id, hist_val); + if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; + } + __syncthreads(); + if (tx < kRadixBins) { + // Inter-warp prefix (only first kHistWarps warp totals matter) + const auto tmp = (lane_id < kRadixWarps) ? smem->warp_sum[lane_id] : 0; + const auto total = warp::reduce_sum(tmp); + const auto inter = warp::reduce_sum(lane_id < warp_id ? tmp : 0); + const auto prefix = inter + warp_inc; // inclusive prefix through this bin + const auto above = total - prefix; // elements in bins ABOVE this one + // 3. Find threshold bin + if (above < topk_remain && above + hist_val >= topk_remain) { + smem->match = {tx, above, topk_remain - above}; + } + } + __syncthreads(); + + const auto [thr, n_above, _] = smem->match; + + // 4. Scatter + if (active) { + if (bin > thr) { + write_pos = num_above + atomicAdd(&smem->counter, 1); + active = false; + } else if (bin < thr) { + active = false; + } else if (round == 3) { + write_pos = K - atomicAdd(&smem->match.equal_count, -1u); + } + // my_bin == thr && round < 3: stay active for next round + } + + topk_remain -= n_above; + if (topk_remain == 0) break; + } + + if (write_pos < K) params.write(write_pos, idx); +} + +} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh new file mode 100644 index 0000000000..73eef555f4 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh @@ -0,0 +1,54 @@ +#pragma once +#include + +#include + +#include + +namespace device::top512 { + +namespace ptx { + +SGL_DEVICE void mbarrier_wait(uint64_t* addr, uint32_t phase) { + while (!cuda::ptx::mbarrier_try_wait_parity(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, addr, phase)) + ; +} + +SGL_DEVICE void mbarrier_init(uint64_t* addr, uint32_t arrives) { + cuda::ptx::mbarrier_init(addr, arrives); +} + +SGL_DEVICE void mbarrier_arrive_expect_tx(uint64_t* addr, uint32_t tx) { + cuda::ptx::mbarrier_arrive_expect_tx(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, cuda::ptx::space_shared, addr, tx); +} + +SGL_DEVICE void mbarrier_arrive(uint64_t* addr) { + cuda::ptx::mbarrier_arrive(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, cuda::ptx::space_shared, addr); +} + +SGL_DEVICE void tma_load(void* dst, const void* src, uint32_t num_bytes, uint64_t* mbar) { + cuda::ptx::cp_async_bulk(cuda::ptx::space_shared, cuda::ptx::space_global, dst, src, num_bytes, mbar); +} + +SGL_DEVICE uint32_t elect_sync() { + uint32_t pred = 0; + asm volatile( + "{\n\t" + ".reg .pred %%px;\n\t" + "elect.sync _|%%px, %1;\n\t" + "@%%px mov.s32 %0, 1;\n\t" + "}" + : "+r"(pred) + : "r"(0xFFFFFFFF)); + return pred; +} + +SGL_DEVICE bool elect_sync_cta(uint32_t tx) { + const auto warp_id = tx / 32; + const auto uniform_warp_id = __shfl_sync(0xFFFFFFFF, warp_id, 0); + return (uniform_warp_id == 0 && elect_sync()); +} + +} // namespace ptx + +} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh new file mode 100644 index 0000000000..77d7361ee8 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh @@ -0,0 +1,302 @@ +#pragma once + +#include +#include +#include + +#include "common.cuh" +#include "ptx.cuh" +#include +#include + +namespace device::top512 { + +template +struct RegisterTopK { + static constexpr uint32_t kHistBits = 12; + static constexpr uint32_t kHistBins = 1 << kHistBits; + static constexpr uint32_t kVecsPerThread = 4; + static constexpr uint32_t kMaxTolerance = 0; + static constexpr uint32_t kMax1PassLength = kVecsPerThread * 4 * kBlockSize; + static constexpr uint32_t kMaxExtraLength = kMax1PassLength; + static constexpr uint32_t kMax2PassLength = kMax1PassLength + kMaxExtraLength; + + struct Smem { + using HistVec = AlignedVector; + alignas(128) uint32_t counter_gt; + alignas(128) uint32_t counter_eq; + uint64_t mbarrier; // for cp.async + MatchBin match; + uint32_t warp_sum[kNumWarps]; + union { + uint32_t histogram[kHistBins]; + HistVec histogram_vec[kBlockSize]; + Tie tie_buffer[kMaxTies]; + }; + alignas(16) float score_buffer[kMaxExtraLength]; + }; + + template + SGL_DEVICE static void + run(const float* scores, // + int32_t* indices, + const uint32_t length, + void* _smem, + const bool use_pdl = false) { + const auto smem = static_cast(_smem); + const auto tx = threadIdx.x; + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + // Initialize shared memory histogram + { + typename Smem::HistVec hist_vec; + hist_vec.fill(0); + smem->histogram_vec[tx] = hist_vec; + if (tx == 0) { + smem->counter_gt = smem->counter_eq = 0; + if constexpr (kIs2Pass) { + ptx::mbarrier_init(&smem->mbarrier, 1); + } + } + __syncthreads(); + } + + if (use_pdl) device::PDLWaitPrimary(); + + // Load scores into registers + Vec4 local[kVecsPerThread]; +#pragma unroll + for (uint32_t v = 0; v < kVecsPerThread; ++v) { + const uint32_t base = (tx + v * kBlockSize) * 4; + if (base >= length) break; + local[v].load(scores, tx + v * kBlockSize); + } + + // Fetch the next chunk of scores + if constexpr (kIs2Pass) { + if (ptx::elect_sync_cta(tx)) { + const auto length_aligned = (length + 3u - kMax1PassLength) & ~3u; + const auto size_bytes = length_aligned * sizeof(float); + ptx::tma_load(smem->score_buffer, scores + kMax1PassLength, size_bytes, &smem->mbarrier); + ptx::mbarrier_arrive_expect_tx(&smem->mbarrier, size_bytes); + } + __syncwarp(); // avoid warp divergence on + } + + // Accumulate histogram via shared-memory atomics +#pragma unroll + for (uint32_t v = 0; v < kVecsPerThread; ++v) { +#pragma unroll + for (uint32_t e = 0; e < 4; ++e) { + if constexpr (!kIs2Pass) { + const uint32_t idx = (tx + v * kBlockSize) * 4 + e; + if (idx >= length) goto LABEL_ACC_FINISH; + } + atomicAdd(&smem->histogram[extract_coarse_bin(local[v][e])], 1); + } + } + if constexpr (kIs2Pass) { + // 16K ~ 32K. `i` is a float4 index + if (lane_id == 0) ptx::mbarrier_wait(&smem->mbarrier, 0); + __syncwarp(); + for (uint32_t i = tx; i + kMax1PassLength < length; i += kBlockSize) { + const auto val = smem->score_buffer[i]; + atomicAdd(&smem->histogram[extract_coarse_bin(val)], 1); + } + } + [[maybe_unused]] LABEL_ACC_FINISH: + __syncthreads(); + + // Phase 2: Exclusive prefix scan -> find threshold bin + { + constexpr uint32_t kItems = kHistBins / kBlockSize; + uint32_t orig[kItems]; + const auto hist_vec = smem->histogram_vec[tx]; + uint32_t tmp_local_sum = 0; + +#pragma unroll + for (uint32_t i = 0; i < kItems; ++i) { + orig[i] = hist_vec[i]; + tmp_local_sum += orig[i]; + } + + const auto warp_inc = warp_inclusive_sum(lane_id, tmp_local_sum); + const auto warp_exc = warp_inc - tmp_local_sum; + if (lane_id == kWarpThreads - 1) { + smem->warp_sum[warp_id] = warp_inc; + } + + __syncthreads(); + + const auto tmp = smem->warp_sum[lane_id]; + // Exactly one bin satisfies: above < K && above + count >= K + uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); + prefix_sum += warp_exc; +#pragma unroll + for (uint32_t i = 0; i < kItems; ++i) { + prefix_sum += orig[i]; + const auto above = length - prefix_sum; + if (above < K && above + orig[i] >= K) { + smem->match = { + .bin = tx * kItems + i, + .above_count = above, + .equal_count = orig[i], + }; + } + } + __syncthreads(); + } + + const auto [thr_bin, num_above, num_equal] = smem->match; + + // Phase 3: Scatter + // Elements strictly above threshold go directly to output. + // Tied elements: simple path admits first-come; tiebreak path collects into tie_buffer. + const bool need_tiebreak = (num_equal + num_above > K + kMaxTolerance); + const auto topk_indices = indices; + const auto tie_buffer = smem->tie_buffer; + +#pragma unroll + for (uint32_t v = 0; v < kVecsPerThread; ++v) { +#pragma unroll + for (uint32_t e = 0; e < 4; ++e) { + const uint32_t idx = (tx + v * kBlockSize) * 4 + e; + if constexpr (!kIs2Pass) { + if (idx >= length) goto LABEL_SCATTER_DONE; + } + const uint32_t bin = extract_coarse_bin(local[v][e]); + if (bin > thr_bin) { + topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx; + } else if (bin == thr_bin) { + const auto pos = atomicAdd(&smem->counter_eq, 1); + if (need_tiebreak) { + if (pos < kMaxTies) { + tie_buffer[pos] = {.idx = idx, .score = local[v][e]}; + } + } else { + if (const auto which = pos + num_above; which < K) { + topk_indices[which] = idx; + } + } + } + } + // prefetch the next scores + if constexpr (kIs2Pass) { + local[v].load(smem->score_buffer, tx + v * kBlockSize); + } + } + + // 16K ~ 32K, already in registers: similar loop as above but read from smem->score_buffer + if constexpr (kIs2Pass) { +#pragma unroll + for (uint32_t v = 0; v < kVecsPerThread; ++v) { +#pragma unroll + for (uint32_t e = 0; e < 4; ++e) { + const uint32_t idx = (tx + v * kBlockSize) * 4 + e + kMax1PassLength; + if (idx >= length) goto LABEL_SCATTER_DONE; + const uint32_t bin = extract_coarse_bin(local[v][e]); + if (bin > thr_bin) { + topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx; + } else if (bin == thr_bin) { + const auto pos = atomicAdd(&smem->counter_eq, 1); + if (need_tiebreak) { + if (pos < kMaxTies) { + tie_buffer[pos] = {.idx = idx, .score = local[v][e]}; + } + } else { + if (const auto which = pos + num_above; which < K) { + topk_indices[which] = idx; + } + } + } + } + } + } + + [[maybe_unused]] LABEL_SCATTER_DONE: + if (!need_tiebreak) return; + + // Phase 4: Tie-breaking within the threshold bin. + // Assume num_ties <= kBlockSize (at most 1 block of ties). + // Each thread takes one tied element, computes its rank (number of + // elements with strictly higher score, breaking exact float ties by + // original index), and writes to output if rank < topk_remain. + __syncthreads(); + static_assert(kMaxTies <= kBlockSize); + + const uint32_t num_ties = min(num_equal, kMaxTies); + const uint32_t topk_remain = K - num_above; + + const auto is_greater = [](const Tie& a, const Tie& b) { + return (a.score > b.score) || (a.score == b.score && a.idx < b.idx); + }; + + if (num_ties <= kWarpThreads) { + static_assert(kWarpThreads <= kNumWarps); + if (lane_id >= num_ties || warp_id >= num_ties) return; // some threads are idle + /// NOTE: use long long to avoid mask overflow when num_ties == 32 + const uint32_t mask = (1ull << num_ties) - 1u; + const auto tie = tie_buffer[lane_id]; + const auto target_tie = tie_buffer[warp_id]; + const bool pred = is_greater(tie, target_tie); + const auto rank = static_cast(__popc(__ballot_sync(mask, pred))); + if (lane_id == 0 && rank < topk_remain) { + topk_indices[num_above + rank] = target_tie.idx; + } + } else if (num_ties <= kWarpThreads * 2) { + // 64 x 64 topk implementation: each thread takes 2 elements + const auto lane_id_1 = lane_id + kWarpThreads; + const auto warp_id_1 = warp_id + kWarpThreads; + const auto invalid = Tie{.idx = 0xFFFFFFFF, .score = -FLT_MAX}; + const auto tie_0 = tie_buffer[lane_id]; + const auto tie_1 = lane_id_1 < num_ties ? tie_buffer[lane_id_1] : invalid; + if (true) { + const auto target = tie_buffer[warp_id]; + const bool pred_0 = is_greater(tie_0, target); + const bool pred_1 = is_greater(tie_1, target); + const auto rank_0 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_0))); + const auto rank_1 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_1))); + const auto rank = rank_0 + rank_1; + if (lane_id == 0 && rank < topk_remain) { + topk_indices[num_above + rank] = target.idx; + } + } + if (warp_id_1 < num_ties) { + const auto target = tie_buffer[warp_id_1]; + const bool pred_0 = is_greater(tie_0, target); + const bool pred_1 = is_greater(tie_1, target); + const auto rank_0 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_0))); + const auto rank_1 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_1))); + const auto rank = rank_0 + rank_1; + if (lane_id == 0 && rank < topk_remain) { + topk_indices[num_above + rank] = target.idx; + } + } + } else { + /// NOTE: Based on my observation, this path is very rarely reached + [[unlikely]]; + // Block-level: each thread reads from tie_buffer in shared memory + for (auto i = warp_id; i < num_ties; i += kNumWarps) { + const auto target_tie = tie_buffer[i]; + uint32_t local_rank = 0; + for (auto j = lane_id; j < num_ties; j += kWarpThreads) { + const auto tie = tie_buffer[j]; + if (is_greater(tie, target_tie)) local_rank++; + } + // sum the rank across the warp + const auto rank = warp::reduce_sum(local_rank); + if (lane_id == 0 && rank < topk_remain) { + topk_indices[num_above + rank] = target_tie.idx; + } + } + } + } + + SGL_DEVICE static void transform(const TransformParams params) { + __syncthreads(); + if (const auto tx = threadIdx.x; tx < K) params.transform(tx); + } +}; + +} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh new file mode 100644 index 0000000000..4462b89a19 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh @@ -0,0 +1,213 @@ +#pragma once + +#include +#include +#include + +#include "common.cuh" +#include "ptx.cuh" +#include +#include + +namespace device::top512 { + +template +struct StreamingTopK { + static constexpr uint32_t kHistBits = 12; + static constexpr uint32_t kHistBins = 1 << kHistBits; + static constexpr uint32_t kRadixBins = 256; + static constexpr uint32_t kElemPerStage = 8; + static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize; + static constexpr uint32_t kNumStages = 2; // double buffer + + static constexpr uint32_t kHistItems = kHistBins / kBlockSize; // 4 + static_assert(kHistItems * kBlockSize == kHistBins); + using HistVec = AlignedVector; + + struct Smem { + uint64_t barrier[2][kNumStages]; + alignas(128) uint32_t counter_gt; + alignas(128) uint32_t counter_eq; + alignas(128) MatchBin match; + alignas(128) uint32_t warp_sum[kNumWarps]; + union { + uint32_t histogram[kHistBins]; + HistVec histogram_vec[kBlockSize]; + Tie tie_buffer[kMaxTies]; + }; + union { + float score_buffer[kNumStages][kSizePerStage]; + TieHandleSmem stage2; // reuse smem for tie handling in phase D + }; + }; + + // --------------------------------------------------------------------------- + // Helpers + // --------------------------------------------------------------------------- + + /// NOTE: length must be 4-aligned since we load 4 floats/thread. Caller should round up. + template + SGL_DEVICE static void issue_tma(const float* scores, uint32_t stage, uint32_t length, Smem* smem) { + const auto buf_idx = stage % kNumStages; + const auto offset = stage * kSizePerStage; + const auto size = min(kSizePerStage, length - offset); + const auto size_bytes = size * sizeof(float); + const auto bar = &smem->barrier[kIsScatter][buf_idx]; + ptx::tma_load(smem->score_buffer[buf_idx], scores + offset, size_bytes, bar); + ptx::mbarrier_arrive_expect_tx(bar, size_bytes); + } + + // --------------------------------------------------------------------------- + // Unified streaming pass. Used for both phase A (kIsScatter=false) and + // phase C (kIsScatter=true). Each buffer is reused across iterations via the + // reuse-arrive trick (same pattern as ClusterTopKImpl::stage1). + // --------------------------------------------------------------------------- + + template + SGL_DEVICE static void stream_pass( + const float* scores, + const uint32_t length, + const uint32_t thr_bin, // ignored when !kIsScatter + int32_t* s_topk_indices, // ignored when !kIsScatter + Smem* smem) { + const auto tx = threadIdx.x; + const auto num_iters = (length + kSizePerStage - 1) / kSizePerStage; + const auto lane_id = tx % kWarpThreads; + + // Initial double-buffer TMA prologue. + const auto length_aligned = (length + 3u) & ~3u; + if (tx == 0) { +#pragma unroll + for (uint32_t i = 0; i < kNumStages; i++) { + if (i >= num_iters) break; + issue_tma(scores, i, length_aligned, smem); + } + } + + for (uint32_t iter = 0; iter < num_iters; iter++) { + const auto buf_idx = iter % kNumStages; + const auto offset = iter * kSizePerStage; + const auto this_size = min(kSizePerStage, length - offset); + + if (lane_id == 1) { + const auto phase_bit = (iter / kNumStages) & 1; + ptx::mbarrier_wait(&smem->barrier[kIsScatter][buf_idx], phase_bit); + } + __syncwarp(); + +#pragma unroll + for (uint32_t i = 0; i < kElemPerStage; i++) { + const auto local_idx = tx + i * kBlockSize; + if (local_idx >= this_size) break; + const auto score = smem->score_buffer[buf_idx][local_idx]; + const auto bin = extract_coarse_bin(score); + if constexpr (kIsScatter) { + const auto global_idx = offset + local_idx; + if (bin > thr_bin) { + const auto pos = atomicAdd(&smem->counter_gt, 1); + if (pos < K) s_topk_indices[pos] = global_idx; + } else if (bin == thr_bin) { + const auto pos = atomicAdd(&smem->counter_eq, 1); + if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score}; + } + } else { + atomicAdd(&smem->histogram[bin], 1); + } + } + + __syncthreads(); + if (tx == 0) { + if (const auto next_iter = iter + kNumStages; next_iter < num_iters) { + issue_tma(scores, next_iter, length_aligned, smem); + } + } + } + } + + // --------------------------------------------------------------------------- + // Phase B: find the threshold bin via a warp-level prefix scan. + // Same structure as SmallTopKImpl's phase 2 (4 bins/thread, warp_sum relay). + // --------------------------------------------------------------------------- + + SGL_DEVICE static void find_threshold(uint32_t length, Smem* smem) { + const auto tx = threadIdx.x; + const auto lane_id = tx % kWarpThreads; + const auto warp_id = tx / kWarpThreads; + + uint32_t orig[kHistItems]; + const auto hist_vec = smem->histogram_vec[tx]; + uint32_t local_sum = 0; +#pragma unroll + for (uint32_t i = 0; i < kHistItems; ++i) { + orig[i] = hist_vec[i]; + local_sum += orig[i]; + } + + const auto warp_inc = warp_inclusive_sum(lane_id, local_sum); + const auto warp_exc = warp_inc - local_sum; + if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; + __syncthreads(); + + const auto tmp = smem->warp_sum[lane_id]; + uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); + prefix_sum += warp_exc; +#pragma unroll + for (uint32_t i = 0; i < kHistItems; ++i) { + prefix_sum += orig[i]; + const auto above = length - prefix_sum; + if (above < K && above + orig[i] >= K) { + smem->match = { + .bin = tx * kHistItems + i, + .above_count = above, + .equal_count = orig[i], + }; + } + } + __syncthreads(); + } + + SGL_DEVICE static void run(const float* scores, const uint32_t length, int32_t* topk_indices, void* _smem) { + const auto smem = static_cast(_smem); + const auto tx = threadIdx.x; + __builtin_assume(tx < kBlockSize); + + // Init histogram, barriers, counters. + { + HistVec zero; + zero.fill(0); + smem->histogram_vec[tx] = zero; + if (tx < 2 * kNumStages) { + const auto base_barrier = &smem->barrier[0][0]; + ptx::mbarrier_init(&base_barrier[tx], 1); + } + if (tx == 0) { + smem->counter_gt = 0; + smem->counter_eq = 0; + } + __syncthreads(); + } + + // Phase A: histogram pass (pipelined TMA stream). + stream_pass(scores, length, 0, nullptr, smem); + + // Phase B: locate threshold bin & re-init barriers + find_threshold(length, smem); + + // Phase C: scatter pass. + stream_pass(scores, length, smem->match.bin, topk_indices, smem); + } + + SGL_DEVICE static void transform(const TransformParams params, void* _smem) { + // Phase D: page-translate above entries, then refine ties. + const auto smem = static_cast(_smem); + const auto tx = threadIdx.x; + const auto num_above = smem->match.above_count; + if (tx < num_above) params.transform(tx); + const auto num_equal = smem->counter_eq; + if (num_above >= K || num_equal == 0) return; + const auto clamped_ties = min(num_equal, kMaxTies); + tie_handle_transform(smem->tie_buffer, clamped_ties, num_above, K, params, &smem->stage2); + } +}; + +} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh new file mode 100644 index 0000000000..e0ce2dc086 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh @@ -0,0 +1,120 @@ +#pragma once +#include + +namespace device::distributed { + +inline constexpr uint32_t kMaxNumGPU = 8; + +struct alignas(128) Semaphore { + public: + constexpr Semaphore() : m_flag(0), m_counter(0) {} + + template + SGL_DEVICE uint32_t get() const { + uint32_t val; + if constexpr (kFence) { + asm volatile("ld.acquire.sys.global.u32 %0, [%1];" : "=r"(val) : "l"(&m_flag)); + } else { + asm volatile("ld.volatile.global.u32 %0, [%1];" : "=r"(val) : "l"(&m_flag)); + } + return val; + } + + template + SGL_DEVICE uint32_t add(uint32_t val) { + uint32_t old_val; + if constexpr (kFence) { + asm volatile("atom.release.sys.global.add.u32 %0, [%1], %2;" : "=r"(old_val) : "l"(&m_flag), "r"(val)); + } else { + asm volatile("atom.global.add.u32 %0, [%1], %2;" : "=r"(old_val) : "l"(&m_flag), "r"(val)); + } + return old_val; + } + + // Only called by the owning GPU - plain load is sufficient + SGL_DEVICE uint32_t get_counter() const { + return m_counter; + } + + // Only called by the owning GPU - plain store is sufficient + SGL_DEVICE void set_counter(uint32_t val) { + m_counter = val; + } + + private: + uint32_t m_flag; + uint32_t m_counter; +}; + +struct PullController { + public: + using SignalType = Semaphore; + + PullController(void** signals, uint32_t num_gpu) { + for (uint32_t i = 0; i < num_gpu; ++i) { + m_signals[i] = static_cast(signals[i]); + } + } + + /// Synchronize all GPUs. + /// When kFence is true, establishes happens-before across GPUs using + /// release/acquire semantics, ensuring prior writes are visible system-wide. + template + SGL_DEVICE void sync(uint32_t rank, uint32_t num_gpu) const { + // For fenced sync: ensure all threads in this block have completed their writes, + // so the signaling thread's release carries them transitively. + static_assert(!(kFence && kStart), "Start stage does not need to wait fence"); + if constexpr (kFence || !kStart) __syncthreads(); + constexpr auto kStage = kStart ? 1 : 2; + const auto warp_id = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + if (lane_id == 0 && warp_id < num_gpu) { + auto& signal = m_signals[warp_id][blockIdx.x]; + signal.add(1); + if (warp_id == rank) { + const auto target = num_gpu * kStage; + /// NOTE: correctness here: + /// - base is only read/updated locally by the owning GPU + const auto base = signal.get_counter(); + while (signal.get() - base < target) + ; + if constexpr (!kStart) { + signal.set_counter(base + target); + } + } + } + if constexpr (kStart) __syncthreads(); + } + + private: + Semaphore* __restrict__ m_signals[kMaxNumGPU]; +}; + +struct PushController { + public: + using SignalType = uint32_t; + static constexpr int64_t kNumStages = 2; + + PushController(void* ptr) : m_local_signal(static_cast(ptr)) {} + + SGL_DEVICE SignalType epoch() const { + return m_local_signal[blockIdx.x]; + } + + SGL_DEVICE void exit() const { + __syncthreads(); + if (threadIdx.x == 0) { + this->exit_unsafe(blockIdx.x); + } + } + + SGL_DEVICE void exit_unsafe(uint32_t which) const { + auto& signal = m_local_signal[which]; + signal = (signal + 1) % kNumStages; + } + + private: + SignalType* m_local_signal; +}; + +} // namespace device::distributed diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh new file mode 100644 index 0000000000..239fac71a1 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh @@ -0,0 +1,354 @@ +#pragma once +#include + +#include +#include + +#include + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace host::distributed { + +using device::distributed::PullController, device::distributed::PushController; + +struct AllReduceData { + constexpr AllReduceData() {} + void* __restrict__ input[device::distributed::kMaxNumGPU]; +}; + +using ExternHandle = tvm::ffi::Array; + +inline ExternHandle to_extern_handle(void* ptr) { + ExternHandle array; + cudaIpcMemHandle_t handle; + RuntimeDeviceCheck(cudaIpcGetMemHandle(&handle, ptr)); + for (size_t i = 0; i < sizeof(handle); ++i) { + array.push_back(handle.reserved[i]); + } + return array; +} + +inline void* from_extern_handle(const ExternHandle& array) { + cudaIpcMemHandle_t handle; + RuntimeCheck(array.size() == sizeof(handle), "Invalid IPC handle size: ", array.size()); + for (size_t i = 0; i < sizeof(handle); ++i) { + handle.reserved[i] = array[i]; + } + void* ptr; + RuntimeDeviceCheck(cudaIpcOpenMemHandle(&ptr, handle, cudaIpcMemLazyEnablePeerAccess)); + return ptr; +} + +struct HandleHash { + std::size_t operator()(const cudaIpcMemHandle_t& handle) const { + return std::hash{}({handle.reserved, sizeof(handle.reserved)}); + } +}; + +struct HandleEqual { + bool operator()(const cudaIpcMemHandle_t& a, const cudaIpcMemHandle_t& b) const { + return std::memcmp(a.reserved, b.reserved, sizeof(a.reserved)) == 0; + } +}; + +/** + * \brief The control plane of the custom all-reduce implementation. + * It manages the internal state and synchronization of the participating GPUs. + */ +struct CustomAllReduceBase : public tvm::ffi::Object { + public: + TVM_FFI_DECLARE_OBJECT_INFO_FINAL("sgl.CustomAllReduce", CustomAllReduceBase, tvm::ffi::Object); + + static constexpr bool _type_mutable = true; + using InputPair = tvm::ffi::Tuple; // (offset, ipc handle) + + CustomAllReduceBase( + uint32_t rank, + uint32_t num_gpu, + uint32_t max_num_cta_pull, + uint32_t max_num_cta_push, + int64_t pull_buffer_size, + int64_t push_buffer_size, + int64_t graph_buffer_count) + : m_pull_buffer_bytes(pull_buffer_size), + m_push_buffer_bytes(push_buffer_size), + m_graph_buffer_count(graph_buffer_count), + m_rank(rank), + m_num_gpu(num_gpu), + m_max_num_cta_pull(max_num_cta_pull), + m_max_num_cta_push(max_num_cta_push), + // default config for pull kernel, can be updated by `configure()` + m_num_cta(max_num_cta_pull), + m_cta_size(256) { + RuntimeCheck(pull_buffer_size % 128 == 0, "Pull buffer size should be aligned to 128 bytes"); + RuntimeCheck(push_buffer_size % 128 == 0, "Push buffer size should be aligned to 128 bytes"); + RuntimeCheck(rank < num_gpu, "Invalid rank: ", rank); + const int64_t kU32Max = static_cast(std::numeric_limits::max()); + const int64_t push_buffer_size_all = push_all_ranks_bytes(); + RuntimeCheck(pull_buffer_size <= kU32Max, "Pull buffer size is too large: ", pull_buffer_size); + RuntimeCheck(push_buffer_size_all <= kU32Max, "Push buffer size is too large: ", push_buffer_size_all); + RuntimeDeviceCheck(cudaMalloc(&m_storage, storage_bytes())); + } + + ExternHandle share_storage() { + return to_extern_handle(m_storage); + } + + tvm::ffi::Array share_graph_inputs() { + tvm::ffi::Array result; + const auto new_inputs_count = registered_count() - m_cum_registered_count; + RuntimeCheck(new_inputs_count >= 0, "Invalid new count: ", new_inputs_count); + result.reserve(new_inputs_count); + std::unordered_map ipc_cache; + const auto get_handle = [&](void* ptr) -> ExternHandle { + const auto it = ipc_cache.find(ptr); + if (it != ipc_cache.end()) return it->second; + const auto handle = to_extern_handle(ptr); + ipc_cache.try_emplace(ptr, handle); + return handle; + }; + for (const auto ptr : std::span(m_graph_capture_inputs).subspan(m_cum_registered_count)) { + // note: must share the base address of each allocation, or we get wrong address + void* base_ptr; + const auto cu_result = cuPointerGetAttribute(&base_ptr, CU_POINTER_ATTRIBUTE_RANGE_START_ADDR, (CUdeviceptr)ptr); + RuntimeCheck(cu_result == CUDA_SUCCESS, "failed to get pointer attr"); + const auto offset = reinterpret_cast(ptr) - reinterpret_cast(base_ptr); + result.push_back(InputPair{offset, get_handle(base_ptr)}); + } + return result; + } + + void post_init(tvm::ffi::Array ipc_storages) { + RuntimeCheck(ipc_storages.size() == m_num_gpu, "Invalid array size: ", ipc_storages.size()); + m_peer_storage.resize(m_num_gpu); + for (const auto i : irange(m_num_gpu)) { + if (i == m_rank) { + m_peer_storage[i] = m_storage; + } else { + m_peer_storage[i] = from_extern_handle(ipc_storages[i]); + } + } + + // set signal buffer to zero + const auto pull_signal = get_pull_signal(m_storage); + RuntimeDeviceCheck(cudaMemset(pull_signal, 0, pull_signal_bytes())); + + // update the pull controller and data pointer + RuntimeCheck(!m_pull_ctrl.has_value(), "Controller is already initialized"); + m_pull_ctrl.emplace(m_peer_storage.data(), m_num_gpu); + AllReduceData data; + for (const auto i : irange(m_num_gpu)) { + data.input[i] = get_pull_buffer(m_peer_storage[i]); + } + const auto default_data_ptr = get_data_ptr(); + RuntimeDeviceCheck(cudaMemcpy(default_data_ptr, &data, sizeof(AllReduceData), cudaMemcpyHostToDevice)); + + // update the push controller and data pointer + RuntimeCheck(!m_push_ctrl.has_value(), "Controller is already initialized"); + const auto push_signal = get_push_signal(m_storage); + RuntimeDeviceCheck(cudaMemset(push_signal, 0, push_signal_bytes())); + m_push_ctrl.emplace(push_signal); + const auto push_buffer = get_push_buffer(m_storage); + RuntimeDeviceCheck(cudaMemset(push_buffer, 0, push_all_ranks_bytes())); + } + + void register_inputs(tvm::ffi::Array> ipc_graph_inputs) { + RuntimeCheck(ipc_graph_inputs.size() == m_num_gpu); + const auto new_registered_count = registered_count() - m_cum_registered_count; + RuntimeCheck(new_registered_count >= 0, "Invalid registered count: ", new_registered_count); + if (new_registered_count == 0) return; // avoid `m_get_data_ptr()` out-of-bounds + std::vector data; + data.resize(new_registered_count); + const auto open_cached = [&](const ExternHandle& h) -> void* { + RuntimeCheck(h.size() == sizeof(cudaIpcMemHandle_t), "Invalid IPC handle size: ", h.size()); + cudaIpcMemHandle_t handle; + for (size_t i = 0; i < sizeof(handle); ++i) + handle.reserved[i] = h[i]; + const auto [it, success] = m_ipc_cache.try_emplace(handle, nullptr); + if (success) { + void* ptr; + RuntimeDeviceCheck(cudaIpcOpenMemHandle(&ptr, handle, cudaIpcMemLazyEnablePeerAccess)); + it->second = ptr; + } + return it->second; + }; + for (const auto i : irange(ipc_graph_inputs.size())) { + const auto& array = ipc_graph_inputs[i]; + RuntimeCheck(int64_t(array.size()) == new_registered_count); + if (i == m_rank) { + for (const auto j : irange(new_registered_count)) { + data[j].input[i] = m_graph_capture_inputs[m_cum_registered_count + j]; + } + } else { + for (const auto j : irange(new_registered_count)) { + /// NOTE: structural binding will cause intern compiler error... + const auto elem = array[j]; + const auto offset = elem.get<0>(); + const auto ipc_handle = elem.get<1>(); + data[j].input[i] = pointer::offset(open_cached(ipc_handle), offset); + } + } + } + + const auto new_registered_bytes = sizeof(AllReduceData) * new_registered_count; + const auto dst_ptr = get_data_ptr(m_cum_registered_count); + m_cum_registered_count += new_registered_count; + RuntimeDeviceCheck(cudaMemcpy(dst_ptr, data.data(), new_registered_bytes, cudaMemcpyHostToDevice)); + } + + void set_cuda_graph_capture(bool enabled) { + m_is_graph_capturing = enabled; + } + + void free_ipc_handles() { + for (const auto& pair : m_ipc_cache) { + host::RuntimeDeviceCheck(cudaIpcCloseMemHandle(pair.second)); + } + m_ipc_cache.clear(); + } + + void free_storage() { + host::RuntimeDeviceCheck(cudaFree(m_storage)); + m_storage = nullptr; + } + + tvm::ffi::Tuple configure_pull(uint32_t num_cta, uint32_t cta_size) { + using host::RuntimeCheck; + const auto min_cta_size = m_num_gpu * device::kWarpThreads; + RuntimeCheck(num_cta > 0 && num_cta <= m_max_num_cta_pull, "Invalid number of CTAs: ", num_cta); + RuntimeCheck(cta_size >= min_cta_size, "Block size must be at least ", min_cta_size); + const auto old_num_cta = m_num_cta; + const auto old_block_size = m_cta_size; + m_num_cta = num_cta; + m_cta_size = cta_size; + return tvm::ffi::Tuple{old_num_cta, old_block_size}; + } + + protected: + AllReduceData* allocate_graph_capture_input(void* data_ptr) { + const auto count = registered_count(); + RuntimeCheck(count < m_graph_buffer_count, "Graph buffer overflow, increase `graph_buffer_count`!"); + m_graph_capture_inputs.push_back(data_ptr); + return get_data_ptr(count); + } + AllReduceData* get_data_ptr(int64_t which = -1) { + const auto count = registered_count(); + RuntimeCheck(which >= -1 && which < count, "Invalid graph buffer index: ", which, ", count: ", count); + const auto start = get_pull_params(m_storage); + return static_cast(start) + (1 + which); + } + int64_t registered_count() const { + return static_cast(m_graph_capture_inputs.size()); + } + int64_t pull_signal_bytes() const { + return _align_bytes(sizeof(PullController::SignalType) * m_max_num_cta_pull); + } + int64_t push_signal_bytes() const { + return _align_bytes(sizeof(PushController::SignalType) * m_max_num_cta_push); + } + int64_t graph_param_bytes() const { + return _align_bytes(sizeof(AllReduceData) * (1 + m_graph_buffer_count)); // 1 for default + } + int64_t push_all_ranks_bytes() const { + return _align_bytes(PushController::kNumStages * m_num_gpu * m_push_buffer_bytes); + } + int64_t storage_bytes() const { + return _get_offset_impl(5); + } + void* get_pull_signal(void* ptr) const { + return pointer::offset(ptr, _get_offset_impl(0)); + } + void* get_push_signal(void* ptr) const { + return pointer::offset(ptr, _get_offset_impl(1)); + } + void* get_pull_params(void* ptr) const { + return pointer::offset(ptr, _get_offset_impl(2)); + } + void* get_pull_buffer(void* ptr) const { + return pointer::offset(ptr, _get_offset_impl(3)); + } + void* get_push_buffer(void* ptr) const { + return pointer::offset(ptr, _get_offset_impl(4)); + } + int64_t _get_offset_impl(int64_t which) const { + // | SignalArray (pull + push) | GraphBuffers (pull params) | Buffers (pull + push) | + const int64_t offset_map[5] = { + /*[0]=*/pull_signal_bytes(), + /*[1]=*/push_signal_bytes(), + /*[2]=*/graph_param_bytes(), + /*[3]=*/m_pull_buffer_bytes, + /*[4]=*/push_all_ranks_bytes(), + }; + RuntimeCheck(which >= 0 && which <= 5, "Invalid offset index: ", which); + return std::accumulate(offset_map, offset_map + which, int64_t(0)); + } + static int64_t _align_bytes(int64_t size) { + return div_ceil(size, 128) * 128; + } + + const int64_t m_pull_buffer_bytes; + const int64_t m_push_buffer_bytes; + const int64_t m_graph_buffer_count; + const uint32_t m_rank; + const uint32_t m_num_gpu; + const uint32_t m_max_num_cta_pull; + const uint32_t m_max_num_cta_push; + // these 2 config should only affect pull kernel + uint32_t m_num_cta; + uint32_t m_cta_size; + // other states + bool m_is_graph_capturing = false; + int64_t m_cum_registered_count = 0; + std::optional m_pull_ctrl; + std::optional m_push_ctrl; + void* m_storage = nullptr; + std::vector m_graph_capture_inputs; + std::vector m_peer_storage; + std::unordered_map m_ipc_cache; +}; + +struct CustomAllReduceRef : public tvm::ffi::ObjectRef { + TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(CustomAllReduceRef, tvm::ffi::ObjectRef, CustomAllReduceBase); +}; + +} // namespace host::distributed + +namespace device::distributed { + +template +SGL_DEVICE auto reduce_impl(AlignedVector (&storage)[M]) -> AlignedVector { + fp32x2_t acc[N] = {}; +#pragma unroll // unroll num gpu + for (uint32_t i = 0; i < M; ++i) { +#pragma unroll // unroll vec + for (uint32_t j = 0; j < N; ++j) { + const auto [x, y] = cast(storage[i][j]); + auto& [x_acc, y_acc] = acc[j]; + x_acc += x; + y_acc += y; + } + } + + AlignedVector result; +#pragma unroll + for (uint32_t j = 0; j < N; ++j) { + result[j] = cast(acc[j]); + } + + return result; +} + +} // namespace device::distributed diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h new file mode 100644 index 0000000000..17d9048d4c --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h @@ -0,0 +1,104 @@ +#pragma once +#include + +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +namespace host::ffi { + +using tvm::ffi::Tensor, tvm::ffi::TensorView, tvm::ffi::ShapeView; + +inline Tensor empty(ShapeView shape, DLDataType dtype, DLDevice device) { + return Tensor::FromEnvAlloc(::TVMFFIEnvTensorAlloc, shape, dtype, device); +} + +inline Tensor empty_like(TensorView tensor) { + return empty(tensor.shape(), tensor.dtype(), tensor.device()); +} + +struct _dummy_deleter { + void operator()(void*) const {} +}; + +// template + +template +struct FromBlobContext { + [[no_unique_address]] Fn deleter; + int64_t dimension; + int64_t* get_shape() { + return reinterpret_cast(this + 1); + } + int64_t* get_stride() { + return this->get_shape() + dimension; + } +}; + +template +inline Tensor from_blob( + void* data, + ShapeView shape, + DLDataType dtype, + DLDevice device, + Fn&& deleter = {}, + std::optional stride = {}, + uint64_t byte_offset = 0) { + using Context = FromBlobContext>; + const auto ndim = shape.size(); + const auto ctx = [&] { + auto ptr = std::malloc(sizeof(Context) + sizeof(int64_t) * ndim * 2); + auto ctx = static_cast(ptr); + std::construct_at(ctx, std::forward(deleter), static_cast(ndim)); + stdr::copy_n(shape.data(), ndim, ctx->get_shape()); + if (stride.has_value()) { + RuntimeCheck(stride->size() == ndim, "Stride ndim mismatch!"); + stdr::copy_n(stride->data(), ndim, ctx->get_stride()); + } else { + int64_t stride_val = 1; + for (const auto i : irange(ndim)) { + const auto j = ndim - 1 - i; + ctx->get_stride()[j] = stride_val; + stride_val *= shape[j]; + } + } + return ctx; + }(); + const auto tensor = DLTensor{ + .data = data, + .device = device, + .ndim = static_cast(ndim), + .dtype = dtype, + .shape = ctx->get_shape(), + .strides = ctx->get_stride(), + .byte_offset = byte_offset, + }; + const auto blob_deleter = [](DLManagedTensor* self) { + auto ctx = static_cast(self->manager_ctx); + ctx->deleter(self->dl_tensor.data); + std::destroy_at(ctx); + std::free(ctx); + }; + auto managed_tensor = DLManagedTensor{tensor, ctx, blob_deleter}; + return Tensor::FromDLPack(&managed_tensor); +} + +template +inline Tensor from_blob_like( + void* data, + TensorView t, + Fn&& deleter = {}, + bool is_contiguous = false, // if override to true, the stride will be ignored + uint64_t byte_offset = 0) { + const auto stride = is_contiguous ? std::nullopt : std::optional{t.strides()}; + return from_blob(data, t.shape(), t.dtype(), t.device(), std::forward(deleter), stride, byte_offset); +} + +} // namespace host::ffi diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh new file mode 100644 index 0000000000..cd024acd46 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh @@ -0,0 +1,168 @@ +#pragma once +#include +#include +#include +#include +#include + +#include +#include + +namespace host::norm { + +/** + * \brief Check if the given configuration is supported. + * \tparam T Element type (only fp16_t/bf16_t is supported) + * \tparam kDim Dimension size (usually hidden size) + */ +template +inline constexpr bool is_config_supported() { + if (!std::is_same_v && !std::is_same_v) return false; + if (kDim <= 256) { + return (kDim == 64 || kDim == 128 || kDim == 256); + } else { + return (kDim % 256 == 0 && kDim <= 8192); + } +} + +/** + * \brief Determine whether to use cta norm based on dimension size. + * TL;DR: use warp norm for dim <= 256, cta norm otherwise. + * \tparam T Element type (fp16_t or bf16_t) + * \tparam kDim Dimension size (usually hidden size) + * \note This function assumes that the configuration is supported. + * \see `is_config_supported` + */ +template +inline constexpr bool should_use_cta() { + static_assert(is_config_supported(), "Unsupported norm configuration"); + return kDim > 256; +} + +/** + * \brief Get the number of threads per CTA for cta norm. + * \tparam T Element type (fp16_t or bf16_t) + * \tparam kDim Dimension size (usually hidden size) + * \return Number of threads per CTA + */ +template +inline constexpr uint32_t get_cta_threads() { + static_assert(should_use_cta()); + return (kDim / 256) * device::kWarpThreads; +} + +} // namespace host::norm + +namespace device::norm { + +namespace details { + +template +SGL_DEVICE AlignedVector apply_norm_impl( + const AlignedVector input, + const AlignedVector weight, + const float eps, + [[maybe_unused]] float* smem_buffer, + [[maybe_unused]] uint32_t num_warps) { + float sum_of_squares = 0.0f; + +#pragma unroll + for (auto i = 0u; i < N; ++i) { + const auto fp32_input = cast(input[i]); + sum_of_squares += fp32_input.x * fp32_input.x; + sum_of_squares += fp32_input.y * fp32_input.y; + } + + sum_of_squares = warp::reduce_sum(sum_of_squares); + float norm_factor; + if constexpr (kUseCTA) { + // need to synchronize across the cta + const auto warp_id = threadIdx.x / kWarpThreads; + smem_buffer[warp_id] = sum_of_squares; + __syncthreads(); + // use the first warp to reduce + if (warp_id == 0) { + const auto tx = threadIdx.x; + const auto local_sum = tx < num_warps ? smem_buffer[tx] : 0.0f; + sum_of_squares = warp::reduce_sum(local_sum); + smem_buffer[32] = math::rsqrt(sum_of_squares / kDim + eps); + } + __syncthreads(); + norm_factor = smem_buffer[32]; + } else { + norm_factor = math::rsqrt(sum_of_squares / kDim + eps); + } + + AlignedVector output; + +#pragma unroll + for (auto i = 0u; i < N; ++i) { + const auto fp32_input = cast(input[i]); + const auto fp32_weight = cast(weight[i]); + output[i] = cast({ + fp32_input.x * norm_factor * fp32_weight.x, + fp32_input.y * norm_factor * fp32_weight.y, + }); + } + + return output; +} + +} // namespace details + +/** + * \brief Apply norm using warp-level implementation. + * \tparam kDim Dimension size + * \tparam T Element type (fp16_t or bf16_t) + * \param input Input vector + * \param weight Weight vector + * \param eps Epsilon value for numerical stability + * \return Normalized output vector + */ +template +SGL_DEVICE T apply_norm_warp(const T& input, const T& weight, float eps) { + static_assert(kDim <= 256, "Warp norm only supports dim <= 256"); + return details::apply_norm_impl(input, weight, eps, nullptr, 0); +} + +/** + * \brief Apply norm using CTA-level implementation. + * \tparam kDim Dimension size + * \tparam T Element type (fp16_t or bf16_t) + * \param input Input vector + * \param weight Weight vector + * \param eps Epsilon value for numerical stability + * \param smem Shared memory buffer + * \param num_warps Number of warps in the CTA + * \return Normalized output vector + */ +template +SGL_DEVICE T apply_norm_cta( + const T& input, const T& weight, float eps, float* smem, uint32_t num_warps = blockDim.x / kWarpThreads) { + static_assert(kDim > 256, "CTA norm only supports dim > 256"); + return details::apply_norm_impl(input, weight, eps, smem, num_warps); +} + +/** + * \brief Storage type for norm operation. + * For warp norm, the storage size depends on kDim. + * For cta norm, the storage size is fixed to 16B. + * We will also pack the input 16-bit floats into 32-bit types + * for faster CUDA core operations. + * + * \tparam T Element type (fp16_t or bf16_t) + * \tparam kDim Dimension size + */ +template +using StorageType = std::conditional_t< // storage type + (kDim > 256), // whether to use cta norm + AlignedVector, 4>, // cta norm storage, fixed to 16B + AlignedVector, kDim / (2 * kWarpThreads)> // warp norm storage + >; + +/** + * \brief Minimum shared memory size (in bytes) required for cta norm. + */ +inline constexpr uint32_t kSmemBufferSize = 33; + +} // namespace device::norm diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh new file mode 100644 index 0000000000..4f9ac48141 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh @@ -0,0 +1,71 @@ +/// \file math.cuh +/// \brief Device-side math helper functions and constants. +/// +/// Provides type-generic wrappers around CUDA math intrinsics by +/// dispatching through `dtype_trait`. All functions are forced-inline +/// device functions. + +#pragma once +#include + +#include + +namespace device::math { + +/// \brief Constant: log2(e) +inline constexpr float log2e = 1.44269504088896340736f; +/// \brief Constant: ln(2) +inline constexpr float loge2 = 0.693147180559945309417f; +/// \brief Maximum representable value for FP8 E4M3 format. +inline constexpr float FP8_E4M3_MAX = 448.0f; +static_assert(log2e * loge2 == 1.0f, "log2e * loge2 must be 1"); + +/// \brief Returns the larger of `a` and `b`. +template +SGL_DEVICE T max(T a, T b) { + return dtype_trait::max(a, b); +} + +/// \brief Returns the smaller of `a` and `b`. +template +SGL_DEVICE T min(T a, T b) { + return dtype_trait::min(a, b); +} + +/// \brief Returns the absolute value of `a`. +template +SGL_DEVICE T abs(T a) { + return dtype_trait::abs(a); +} + +/// \brief Returns the square root of `a`. +template +SGL_DEVICE T sqrt(T a) { + return dtype_trait::sqrt(a); +} + +/// \brief Returns the reciprocal square root of `a` (i.e. 1 / sqrt(a)). +template +SGL_DEVICE T rsqrt(T a) { + return dtype_trait::rsqrt(a); +} + +/// \brief Returns e^a. +template +SGL_DEVICE T exp(T a) { + return dtype_trait::exp(a); +} + +/// \brief Returns sin(a). +template +SGL_DEVICE T sin(T a) { + return dtype_trait::sin(a); +} + +/// \brief Returns cos(a). +template +SGL_DEVICE T cos(T a) { + return dtype_trait::cos(a); +} + +} // namespace device::math diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh new file mode 100644 index 0000000000..4ea722a3fe --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh @@ -0,0 +1,86 @@ +/// \file runtime.cuh +/// \brief Host-side CUDA runtime query helpers. +/// +/// Thin wrappers around CUDA occupancy and device-property APIs with +/// automatic error checking via `RuntimeDeviceCheck`. + +#pragma once + +#include + +#include +#include +#ifndef USE_ROCM +#include +#else +#include +#ifndef cudaOccupancyMaxActiveBlocksPerMultiprocessor +#define cudaOccupancyMaxActiveBlocksPerMultiprocessor hipOccupancyMaxActiveBlocksPerMultiprocessor +#endif +#ifndef cudaDeviceGetAttribute +#define cudaDeviceGetAttribute hipDeviceGetAttribute +#endif +#ifndef cudaDevAttrMultiProcessorCount +#define cudaDevAttrMultiProcessorCount hipDeviceAttributeMultiprocessorCount +#endif +#ifndef cudaDevAttrComputeCapabilityMajor +#define cudaDevAttrComputeCapabilityMajor hipDeviceAttributeComputeCapabilityMajor +#endif +#ifndef cudaRuntimeGetVersion +#define cudaRuntimeGetVersion hipRuntimeGetVersion +#endif +#ifndef cudaOccupancyAvailableDynamicSMemPerBlock +inline hipError_t +cudaOccupancyAvailableDynamicSMemPerBlock(std::size_t* smem, const void* func, int num_blocks, int block_size) { + // HIP does not expose this directly; return max shared mem as conservative estimate + hipDeviceProp_t prop; + int device; + hipGetDevice(&device); + hipGetDeviceProperties(&prop, device); + *smem = prop.sharedMemPerBlock; + return hipSuccess; +} +#endif +#endif + +namespace host::runtime { + +// Return the maximum number of active blocks per SM for the given kernel +template +inline auto get_blocks_per_sm(T&& kernel, int32_t block_dim, std::size_t dynamic_smem = 0) -> uint32_t { + int num_blocks_per_sm = 0; + RuntimeDeviceCheck( + cudaOccupancyMaxActiveBlocksPerMultiprocessor(&num_blocks_per_sm, kernel, block_dim, dynamic_smem)); + return static_cast(num_blocks_per_sm); +} + +// Return the number of SMs for the given device +inline auto get_sm_count(int device_id) -> uint32_t { + int sm_count; + RuntimeDeviceCheck(cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, device_id)); + return static_cast(sm_count); +} + +// Return the Major compute capability for the given device +inline auto get_cc_major(int device_id) -> int { + int cc_major; + RuntimeDeviceCheck(cudaDeviceGetAttribute(&cc_major, cudaDevAttrComputeCapabilityMajor, device_id)); + return cc_major; +} + +// Return the runtime version +inline auto get_runtime_version() -> int { + int runtime_version; + RuntimeDeviceCheck(cudaRuntimeGetVersion(&runtime_version)); + return runtime_version; +} + +// Return the maximum dynamic shared memory per block for the given kernel +template +inline auto get_available_dynamic_smem_per_block(T&& kernel, int num_blocks, int block_size) -> std::size_t { + std::size_t smem_size; + RuntimeDeviceCheck(cudaOccupancyAvailableDynamicSMemPerBlock(&smem_size, kernel, num_blocks, block_size)); + return smem_size; +} + +} // namespace host::runtime diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp b/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp new file mode 100644 index 0000000000..d229d3a975 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp @@ -0,0 +1,334 @@ +#pragma once + +#include +#include +#ifndef __CUDACC__ +#include +#endif + +namespace host { + +// +// ScalarType can represent a wide range of floating point and integer types, +// in particular it can be used to represent sub-byte data types (something +// that torch.dtype currently does not support). +// +// The type definitions on the Python side can be found in: vllm/scalar_type.py +// these type definitions should be kept up to date with any Python API changes +// here. +// +class ScalarType { + public: + enum NanRepr : uint8_t { + NAN_NONE = 0, // nans are not supported + NAN_IEEE_754 = 1, // nans are: exp all 1s, mantissa not all 0s + NAN_EXTD_RANGE_MAX_MIN = 2, // nans are: exp all 1s, mantissa all 1s + + NAN_REPR_ID_MAX + }; + + constexpr ScalarType( + uint8_t exponent, + uint8_t mantissa, + bool signed_, + int32_t bias, + bool finite_values_only = false, + NanRepr nan_repr = NAN_IEEE_754) + : exponent(exponent), + mantissa(mantissa), + signed_(signed_), + bias(bias), + finite_values_only(finite_values_only), + nan_repr(nan_repr) {}; + + static constexpr ScalarType int_(uint8_t size_bits, int32_t bias = 0) { + return ScalarType(0, size_bits - 1, true, bias); + } + + static constexpr ScalarType uint(uint8_t size_bits, int32_t bias = 0) { + return ScalarType(0, size_bits, false, bias); + } + + // IEEE 754 compliant floating point type + static constexpr ScalarType float_IEEE754(uint8_t exponent, uint8_t mantissa) { + assert(mantissa > 0 && exponent > 0); + return ScalarType(exponent, mantissa, true, 0, false, NAN_IEEE_754); + } + + // IEEE 754 non-compliant floating point type + static constexpr ScalarType float_(uint8_t exponent, uint8_t mantissa, bool finite_values_only, NanRepr nan_repr) { + assert(nan_repr < NAN_REPR_ID_MAX); + assert(mantissa > 0 && exponent > 0); + assert(nan_repr != NAN_IEEE_754); + return ScalarType(exponent, mantissa, true, 0, finite_values_only, nan_repr); + } + + uint8_t const exponent; // size of the exponent field (0 for integer types) + uint8_t const mantissa; // size of the mantissa field (size of the integer + // excluding the sign bit for integer types) + bool const signed_; // flag if the type supports negative numbers (i.e. has a + // sign bit) + int32_t const bias; // stored values equal value + bias, + // used for quantized type + + // Extra Floating point info + bool const finite_values_only; // i.e. no +/-inf if true + NanRepr const nan_repr; // how NaNs are represented + // (not applicable for integer types) + + using Id = int64_t; + + private: + // Field size in id + template + static constexpr size_t member_id_field_width() { + using T = std::decay_t; + return std::is_same_v ? 1 : sizeof(T) * 8; + } + + template + static constexpr auto reduce_members_helper(Fn f, Init val, Member member, Rest... rest) { + auto new_val = f(val, member); + if constexpr (sizeof...(rest) > 0) { + return reduce_members_helper(f, new_val, rest...); + } else { + return new_val; + }; + } + + template + constexpr auto reduce_members(Fn f, Init init) const { + // Should be in constructor order for `from_id` + return reduce_members_helper(f, init, exponent, mantissa, signed_, bias, finite_values_only, nan_repr); + }; + + template + static constexpr auto reduce_member_types(Fn f, Init init) { + constexpr auto dummy_type = ScalarType(0, 0, false, 0, false, NAN_NONE); + return dummy_type.reduce_members(f, init); + }; + + static constexpr auto id_size_bits() { + return reduce_member_types( + [](int acc, auto member) -> int { return acc + member_id_field_width(); }, 0); + } + + public: + // unique id for this scalar type that can be computed at compile time for + // c++17 template specialization this is not needed once we migrate to + // c++20 and can pass literal classes as template parameters + constexpr Id id() const { + static_assert(id_size_bits() <= sizeof(Id) * 8, "ScalarType id is too large to be stored"); + + auto or_and_advance = [](std::pair result, auto member) -> std::pair { + auto [id, bit_offset] = result; + auto constexpr bits = member_id_field_width(); + return {id | (int64_t(member) & ((uint64_t(1) << bits) - 1)) << bit_offset, bit_offset + bits}; + }; + return reduce_members(or_and_advance, std::pair{}).first; + } + + // create a ScalarType from an id, for c++17 template specialization, + // this is not needed once we migrate to c++20 and can pass literal + // classes as template parameters + static constexpr ScalarType from_id(Id id) { + auto extract_and_advance = [id](auto result, auto member) { + using T = decltype(member); + auto [tuple, bit_offset] = result; + auto constexpr bits = member_id_field_width(); + auto extracted_val = static_cast((int64_t(id) >> bit_offset) & ((uint64_t(1) << bits) - 1)); + auto new_tuple = std::tuple_cat(tuple, std::make_tuple(extracted_val)); + return std::pair{new_tuple, bit_offset + bits}; + }; + + auto [tuple_args, _] = reduce_member_types(extract_and_advance, std::pair, int>{}); + return std::apply([](auto... args) { return ScalarType(args...); }, tuple_args); + } + + constexpr int64_t size_bits() const { + return mantissa + exponent + is_signed(); + } + constexpr bool is_signed() const { + return signed_; + } + constexpr bool is_integer() const { + return exponent == 0; + } + constexpr bool is_floating_point() const { + return exponent > 0; + } + constexpr bool is_ieee_754() const { + return is_floating_point() && finite_values_only == false && nan_repr == NAN_IEEE_754; + } + constexpr bool has_nans() const { + return is_floating_point() && nan_repr != NAN_NONE; + } + constexpr bool has_infs() const { + return is_floating_point() && finite_values_only == false; + } + constexpr bool has_bias() const { + return bias != 0; + } + +#ifndef __CUDACC__ + private: + double _floating_point_max() const { + assert(mantissa <= 52 && exponent <= 11); + + uint64_t max_mantissa = (uint64_t(1) << mantissa) - 1; + if (nan_repr == NAN_EXTD_RANGE_MAX_MIN) { + max_mantissa -= 1; + } + + uint64_t max_exponent = (uint64_t(1) << exponent) - 2; + if (nan_repr == NAN_EXTD_RANGE_MAX_MIN || nan_repr == NAN_NONE) { + assert(exponent < 11); + max_exponent += 1; + } + + // adjust the exponent to match that of a double + // for now we assume the exponent bias is the standard 2^(e-1) -1, (where e + // is the exponent bits), there is some precedent for non-standard biases, + // example `float8_e4m3b11fnuz` here: https://github.com/jax-ml/ml_dtypes + // but to avoid premature over complication we are just assuming the + // standard exponent bias until there is a need to support non-standard + // biases + uint64_t exponent_bias = (uint64_t(1) << (exponent - 1)) - 1; + uint64_t exponent_bias_double = (uint64_t(1) << 10) - 1; // double e = 11 + + uint64_t max_exponent_double = max_exponent - exponent_bias + exponent_bias_double; + + // shift the mantissa into the position for a double and + // the exponent + uint64_t double_raw = (max_mantissa << (52 - mantissa)) | (max_exponent_double << 52); + + return *reinterpret_cast(&double_raw); + } + + constexpr std::variant _raw_max() const { + if (is_floating_point()) { + return {_floating_point_max()}; + } else { + assert(size_bits() < 64 || (size_bits() == 64 && is_signed())); + return {(int64_t(1) << mantissa) - 1}; + } + } + + constexpr std::variant _raw_min() const { + if (is_floating_point()) { + assert(is_signed()); + constexpr uint64_t sign_bit_double = (uint64_t(1) << 63); + + double max = _floating_point_max(); + uint64_t max_raw = *reinterpret_cast(&max); + uint64_t min_raw = max_raw | sign_bit_double; + return {*reinterpret_cast(&min_raw)}; + } else { + assert(!is_signed() || size_bits() <= 64); + if (is_signed()) { + // set the top bit to 1 (i.e. INT64_MIN) and the rest to 0 + // then perform an arithmetic shift right to set all the bits above + // (size_bits() - 1) to 1 + return {INT64_MIN >> (64 - size_bits())}; + } else { + return {int64_t(0)}; + } + } + } + + public: + // Max representable value for this scalar type. + // (accounting for bias if there is one) + constexpr std::variant max() const { + return std::visit([this](auto x) -> std::variant { return {x - bias}; }, _raw_max()); + } + + // Min representable value for this scalar type. + // (accounting for bias if there is one) + constexpr std::variant min() const { + return std::visit([this](auto x) -> std::variant { return {x - bias}; }, _raw_min()); + } +#endif // __CUDACC__ + + public: + std::string str() const { + /* naming generally follows: https://github.com/jax-ml/ml_dtypes + * for floating point types (leading f) the scheme is: + * `float_em[flags]` + * flags: + * - no-flags: means it follows IEEE 754 conventions + * - f: means finite values only (no infinities) + * - n: means nans are supported (non-standard encoding) + * for integer types the scheme is: + * `[u]int[b]` + * - if bias is not present it means its zero + */ + if (is_floating_point()) { + auto ret = + "float" + std::to_string(size_bits()) + "_e" + std::to_string(exponent) + "m" + std::to_string(mantissa); + if (!is_ieee_754()) { + if (finite_values_only) { + ret += "f"; + } + if (nan_repr != NAN_NONE) { + ret += "n"; + } + } + return ret; + } else { + auto ret = ((is_signed()) ? "int" : "uint") + std::to_string(size_bits()); + if (has_bias()) { + ret += "b" + std::to_string(bias); + } + return ret; + } + } + + constexpr bool operator==(ScalarType const& other) const { + return mantissa == other.mantissa && exponent == other.exponent && bias == other.bias && signed_ == other.signed_ && + finite_values_only == other.finite_values_only && nan_repr == other.nan_repr; + } +}; + +using ScalarTypeId = ScalarType::Id; + +// "rust style" names generally following: +// https://github.com/pytorch/pytorch/blob/6d9f74f0af54751311f0dd71f7e5c01a93260ab3/torch/csrc/api/include/torch/types.h#L60-L70 +static inline constexpr auto kS4 = ScalarType::int_(4); +static inline constexpr auto kU4 = ScalarType::uint(4); +static inline constexpr auto kU4B8 = ScalarType::uint(4, 8); +static inline constexpr auto kS8 = ScalarType::int_(8); +static inline constexpr auto kU8 = ScalarType::uint(8); +static inline constexpr auto kU8B128 = ScalarType::uint(8, 128); + +static inline constexpr auto kFE2M1f = ScalarType::float_(2, 1, true, ScalarType::NAN_NONE); +static inline constexpr auto kFE3M2f = ScalarType::float_(3, 2, true, ScalarType::NAN_NONE); +static inline constexpr auto kFE4M3fn = ScalarType::float_(4, 3, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN); +static inline constexpr auto kFE8M0fnu = ScalarType(8, 0, false, 0, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN); +static inline constexpr auto kFE5M2 = ScalarType::float_IEEE754(5, 2); +static inline constexpr auto kFE8M7 = ScalarType::float_IEEE754(8, 7); +static inline constexpr auto kFE5M10 = ScalarType::float_IEEE754(5, 10); + +// Fixed width style names, generally following: +// https://github.com/pytorch/pytorch/blob/6d9f74f0af54751311f0dd71f7e5c01a93260ab3/torch/csrc/api/include/torch/types.h#L47-L57 +static inline constexpr auto kInt4 = kS4; +static inline constexpr auto kUint4 = kU4; +static inline constexpr auto kUint4b8 = kU4B8; +static inline constexpr auto kInt8 = kS8; +static inline constexpr auto kUint8 = kU8; +static inline constexpr auto kUint8b128 = kU8B128; + +static inline constexpr auto kFloat4_e2m1f = kFE2M1f; +static inline constexpr auto kFloat6_e3m2f = kFE3M2f; +static inline constexpr auto kFloat8_e4m3fn = kFE4M3fn; +static inline constexpr auto kFloat8_e5m2 = kFE5M2; +static inline constexpr auto kFloat16_e8m7 = kFE8M7; +static inline constexpr auto kFloat16_e5m10 = kFE5M10; + +// colloquial names +static inline constexpr auto kHalf = kFE5M10; +static inline constexpr auto kFloat16 = kHalf; +static inline constexpr auto kBFloat16 = kFE8M7; + +static inline constexpr auto kFloat16Id = kFloat16.id(); +} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h new file mode 100644 index 0000000000..7c9fd52131 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h @@ -0,0 +1,40 @@ +/// \file source_location.h +/// \brief Portable `source_location` wrapper. +/// +/// Uses `std::source_location` when available (C++20), otherwise falls +/// back to a minimal stub that returns empty/zero values. + +#pragma once +#include + +/// NOTE: fallback to a minimal source_location implementation +#if defined(__cpp_lib_source_location) +#include + +using source_location_t = std::source_location; + +#else + +struct source_location_fallback { + public: + static constexpr source_location_fallback current() noexcept { + return source_location_fallback{}; + } + constexpr source_location_fallback() noexcept = default; + constexpr unsigned line() const noexcept { + return 0; + } + constexpr unsigned column() const noexcept { + return 0; + } + constexpr const char* file_name() const noexcept { + return ""; + } + constexpr const char* function_name() const noexcept { + return ""; + } +}; + +using source_location_t = source_location_fallback; + +#endif diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h new file mode 100644 index 0000000000..1ae9233a61 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h @@ -0,0 +1,605 @@ +/// \file tensor.h +/// \brief Tensor validation and symbolic matching utilities. +/// +/// Provides the `TensorMatcher` fluent API for validating tensor shapes, +/// strides, dtypes, and devices at kernel entry points, along with +/// `SymbolicSize`, `SymbolicDType`, and `SymbolicDevice` for capturing +/// and cross-checking tensor metadata across multiple tensors. +/// +/// See the "Tensor Checking" section in the JIT kernel dev guide for +/// usage examples. + +#pragma once +#include + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#ifdef __CUDACC__ +#include +#elif defined(__HIPCC__) +#include +#endif + +namespace host { + +namespace details { + +inline constexpr auto kAnyDeviceID = -1; +inline constexpr auto kAnySize = static_cast(-1); +inline constexpr auto kNullSize = static_cast(-1); +inline constexpr auto kNullDType = static_cast(18u); +inline constexpr auto kNullDevice = static_cast(-1); + +struct SizeRef; +struct DTypeRef; +struct DeviceRef; + +template +struct _dtype_trait {}; + +template +struct _dtype_trait { + inline static constexpr DLDataType value = { + .code = std::is_signed_v ? DLDataTypeCode::kDLInt : DLDataTypeCode::kDLUInt, + .bits = static_cast(sizeof(T) * 8), + .lanes = 1}; +}; + +template +struct _dtype_trait { + inline static constexpr DLDataType value = { + .code = DLDataTypeCode::kDLFloat, .bits = static_cast(sizeof(T) * 8), .lanes = 1}; +}; + +#ifdef __CUDACC__ +template <> +struct _dtype_trait { + inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat, .bits = 16, .lanes = 1}; +}; +template <> +struct _dtype_trait { + inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1}; +}; +template <> +struct _dtype_trait { + inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat8_e4m3fn, .bits = 8, .lanes = 1}; +}; +#elif defined(__HIPCC__) +template <> +struct _dtype_trait { + inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat, .bits = 16, .lanes = 1}; +}; +template <> +struct _dtype_trait { + inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1}; +}; +#endif + +template +struct _device_trait { + inline static constexpr DLDevice value = {.device_type = Code, .device_id = kAnyDeviceID}; +}; + +template +inline constexpr auto kDTypeList = std::array{_dtype_trait::value...}; + +template +inline constexpr auto kDeviceList = std::array{_device_trait::value...}; + +template +struct PrintAbleSpan { + explicit PrintAbleSpan(std::span data) : data(data) {} + std::span data; +}; + +// define DLDataType comparison and printing in root namespace +inline constexpr auto kDeviceStringMap = [] { + constexpr auto map = std::array, 16>{ + std::pair{DLDeviceType::kDLCPU, "cpu"}, + std::pair{DLDeviceType::kDLCUDA, "cuda"}, + std::pair{DLDeviceType::kDLCUDAHost, "cuda_host"}, + std::pair{DLDeviceType::kDLOpenCL, "opencl"}, + std::pair{DLDeviceType::kDLVulkan, "vulkan"}, + std::pair{DLDeviceType::kDLMetal, "metal"}, + std::pair{DLDeviceType::kDLVPI, "vpi"}, + std::pair{DLDeviceType::kDLROCM, "rocm"}, + std::pair{DLDeviceType::kDLROCMHost, "rocm_host"}, + std::pair{DLDeviceType::kDLExtDev, "ext_dev"}, + std::pair{DLDeviceType::kDLCUDAManaged, "cuda_managed"}, + std::pair{DLDeviceType::kDLOneAPI, "oneapi"}, + std::pair{DLDeviceType::kDLWebGPU, "webgpu"}, + std::pair{DLDeviceType::kDLHexagon, "hexagon"}, + std::pair{DLDeviceType::kDLMAIA, "maia"}, + std::pair{DLDeviceType::kDLTrn, "trn"}, + }; + constexpr auto max_type = stdr::max(map | stdv::keys); + auto result = std::array{}; + for (const auto& [code, name] : map) { + result[static_cast(code)] = name; + } + return result; +}(); + +struct PrintableDevice { + DLDevice device; +}; + +inline auto& operator<<(std::ostream& os, DLDevice device) { + const auto& mapping = kDeviceStringMap; + const auto entry = static_cast(device.device_type); + RuntimeCheck(entry < mapping.size()); + const auto name = mapping[entry]; + RuntimeCheck(!name.empty(), "Unknown device: ", int(device.device_type)); + os << name; + if (device.device_id != kAnyDeviceID && device.device_type != DLDeviceType::kDLCPU) { + os << ":" << device.device_id; + } + return os; +} + +inline auto& operator<<(std::ostream& os, PrintableDevice pd) { + return os << pd.device; +} + +template +inline auto& operator<<(std::ostream& os, PrintAbleSpan span) { + os << "["; + for (const auto i : irange(span.data.size())) { + if (i > 0) { + os << ", "; + } + os << span.data[i]; + } + os << "]"; + return os; +} + +} // namespace details + +/// \brief Check whether `dtype` matches the DLDataType for C++ type `T`. +template +inline bool is_type(DLDataType dtype) { + return dtype == details::_dtype_trait::value; +} + +/** + * \brief A symbolic dimension size that can be bound once and + * verified across multiple tensors. + * + * Create with an optional annotation string for error messages: + * \code + * auto N = SymbolicSize{"num_tokens"}; + * \endcode + * + * Call `verify()` during tensor matching to either bind the first + * observed value or check subsequent values match. Call `unwrap()` + * to retrieve the bound value (panics if unset). + */ +struct SymbolicSize { + public: + SymbolicSize(std::string_view annotation = {}) : m_value(details::kNullSize), m_annotation(annotation) {} + SymbolicSize(const SymbolicSize&) = delete; + SymbolicSize& operator=(const SymbolicSize&) = delete; + + auto get_name() const -> std::string_view { + return m_annotation; + } + + auto set_value(int64_t value) -> void { + RuntimeCheck(!this->has_value(), "Size value already set"); + m_value = value; + } + + auto has_value() const -> bool { + return m_value != details::kNullSize; + } + + auto get_value() const -> std::optional { + return this->has_value() ? std::optional{m_value} : std::nullopt; + } + + auto unwrap(DebugInfo info = {}) const -> int64_t { + RuntimeCheck(info, this->has_value(), "Size value is not set"); + return m_value; + } + + auto verify(int64_t value, const char* prefix, int64_t dim) -> void { + if (this->has_value()) { + if (m_value != value) { + [[unlikely]]; + Panic("Size mismatch for ", m_name_str(prefix, dim), ": expected ", m_value, " but got ", value); + } + } else { + this->set_value(value); + } + } + + auto value_or_name(const char* prefix, int64_t dim) const -> std::string { + if (const auto value = this->get_value()) { + return std::to_string(*value); + } else { + return m_name_str(prefix, dim); + } + } + + private: + auto m_name_str(const char* prefix, int64_t dim) const -> std::string { + std::ostringstream os; + os << prefix << '#' << dim; + if (!m_annotation.empty()) os << "('" << m_annotation << "')"; + return std::move(os).str(); + } + + std::int64_t m_value; + std::string_view m_annotation; +}; + +inline auto operator==(DLDevice lhs, DLDevice rhs) -> bool { + return lhs.device_type == rhs.device_type && lhs.device_id == rhs.device_id; +} + +/** + * \brief A symbolic data type that can be constrained and verified. + * + * Optionally restrict allowed types via `set_options()`. + * Use `verify()` to bind/check the dtype, and `unwrap()` to retrieve it. + */ +struct SymbolicDType { + public: + SymbolicDType() : m_value({details::kNullDType, 0, 0}) {} + SymbolicDType(const SymbolicDType&) = delete; + SymbolicDType& operator=(const SymbolicDType&) = delete; + + auto set_value(DLDataType value) -> void { + RuntimeCheck(!this->has_value(), "Dtype value already set"); + RuntimeCheck( + m_check(value), "Dtype value [", value, "] not in the allowed options: ", details::PrintAbleSpan{m_options}); + m_value = value; + } + + auto has_value() const -> bool { + return m_value.code != details::kNullDType; + } + + auto get_value() const -> std::optional { + return this->has_value() ? std::optional{m_value} : std::nullopt; + } + + auto unwrap(DebugInfo info = {}) const -> DLDataType { + RuntimeCheck(info, this->has_value(), "Dtype value is not set"); + return m_value; + } + + auto set_options(std::span options) -> void { + m_options = options; + } + + template + auto set_options() -> void { + m_options = details::kDTypeList; + } + + auto verify(DLDataType dtype) -> void { + if (this->has_value()) { + RuntimeCheck(m_value == dtype, "DType mismatch: expected ", m_value, " but got ", dtype); + } else { + this->set_value(dtype); + } + } + + template + auto is_type() const -> bool { + return ::host::is_type(m_value); + } + + private: + auto m_check(DLDataType value) const -> bool { + return stdr::empty(m_options) || (stdr::find(m_options, value) != stdr::end(m_options)); + } + + std::span m_options; + DLDataType m_value; +}; + +/** + * \brief A symbolic device that can be constrained and verified. + * + * Optionally restrict allowed device types via + * `set_options()`. The device id can be wildcarded. + */ +struct SymbolicDevice { + public: + SymbolicDevice() : m_value({details::kNullDevice, details::kAnyDeviceID}) {} + SymbolicDevice(const SymbolicDevice&) = delete; + SymbolicDevice& operator=(const SymbolicDevice&) = delete; + + auto set_value(DLDevice value) -> void { + RuntimeCheck(!this->has_value(), "Device value already set"); + RuntimeCheck( + m_check(value), + "Device value [", + details::PrintableDevice{value}, + "] not in the allowed options: ", + details::PrintAbleSpan{m_options}); + m_value = value; + } + + auto has_value() const -> bool { + return m_value.device_type != details::kNullDevice; + } + + auto get_value() const -> std::optional { + return this->has_value() ? std::optional{m_value} : std::nullopt; + } + + auto unwrap(DebugInfo info = {}) const -> DLDevice { + RuntimeCheck(info, this->has_value(), "Device value is not set"); + return m_value; + } + + auto set_options(std::span options) -> void { + m_options = options; + } + + template + auto set_options() -> void { + m_options = details::kDeviceList; + } + + auto verify(DLDevice device) -> void { + if (this->has_value()) { + RuntimeCheck( + m_value == device, + "Device mismatch: expected ", + details::PrintableDevice{m_value}, + " but got ", + details::PrintableDevice{device}); + } else { + this->set_value(device); + } + } + + private: + auto m_check(DLDevice value) const -> bool { + return stdr::empty(m_options) || (stdr::any_of(m_options, [value](const DLDevice& opt) { + // device type must exactly match + if (opt.device_type != value.device_type) return false; + // device id can be wildcarded + return opt.device_id == details::kAnyDeviceID || opt.device_id == value.device_id; + })); + } + + std::span m_options; + DLDevice m_value; +}; + +namespace details { + +template +struct BaseRef { + public: + BaseRef(const BaseRef&) = delete; + BaseRef& operator=(const BaseRef&) = delete; + + auto operator->() const -> T* { + return m_ref; + } + auto operator*() const -> T& { + return *m_ref; + } + auto rebind(T& other) -> void { + m_ref = &other; + } + + explicit BaseRef() : m_ref(&m_cache), m_cache() {} + BaseRef(T& size) : m_ref(&size), m_cache() {} + + private: + T* m_ref; + T m_cache; +}; + +struct SizeRef : BaseRef { + using BaseRef::BaseRef; + SizeRef(int64_t value) { + if (value != kAnySize) { + (**this).set_value(value); + } else { + // otherwise, we can match any size + } + } +}; + +struct DTypeRef : BaseRef { + using BaseRef::BaseRef; + DTypeRef(DLDataType options) { + (**this).set_value(options); + } + DTypeRef(std::initializer_list options) { + (**this).set_options(options); + } + DTypeRef(std::span options) { + (**this).set_options(options); + } +}; + +struct DeviceRef : BaseRef { + using BaseRef::BaseRef; + DeviceRef(DLDevice options) { + (**this).set_value(options); + } + DeviceRef(std::initializer_list options) { + (**this).set_options(options); + } + DeviceRef(std::span options) { + (**this).set_options(options); + } +}; + +} // namespace details + +/** + * \brief Fluent API for validating tensor shape, strides, dtype, and device. + * + * Construct with the expected shape (using `SymbolicSize` or literal + * integers), chain `.with_strides()`, `.with_dtype<...>()`, and + * `.with_device<...>()`, then call `.verify(tensor)`. + * + * Example: + * \code + * auto N = SymbolicSize{"N"}; + * TensorMatcher({N, 128}) + * .with_dtype() + * .with_device() + * .verify(input_tensor); + * \endcode + * + * \note `TensorMatcher` is a move-only temporary. Do not store in a variable. + */ +struct TensorMatcher { + private: + using SizeRef = details::SizeRef; + using DTypeRef = details::DTypeRef; + using DeviceRef = details::DeviceRef; + + public: + TensorMatcher(const TensorMatcher&) = delete; + TensorMatcher& operator=(const TensorMatcher&) = delete; + + explicit TensorMatcher(std::initializer_list shape) : m_shape(shape), m_strides(), m_dtype() {} + + auto with_strides(std::initializer_list strides) && -> TensorMatcher&& { + // no partial update allowed + RuntimeCheck(m_strides.size() == 0, "Strides already specified"); + RuntimeCheck(m_shape.size() == strides.size(), "Strides size must match shape size"); + m_strides = strides; + return std::move(*this); + } + + template + auto with_dtype(DTypeRef&& dtype) && -> TensorMatcher&& { + m_init_dtype(); + m_dtype.rebind(*dtype); + m_dtype->set_options(); + return std::move(*this); + } + + template + auto with_dtype() && -> TensorMatcher&& { + static_assert(sizeof...(Ts) > 0, "At least one dtype option must be specified"); + m_init_dtype(); + m_dtype->set_options(); + return std::move(*this); + } + + template + auto with_device(DeviceRef&& device) && -> TensorMatcher&& { + m_init_device(); + m_device.rebind(*device); + m_device->set_options(); + return std::move(*this); + } + + template + auto with_device() && -> TensorMatcher&& { + static_assert(sizeof...(Codes) > 0, "At least one device option must be specified"); + m_init_device(); + m_device->set_options(); + return std::move(*this); + } + + // once we start verification, we cannot modify anymore + auto verify(tvm::ffi::TensorView view, DebugInfo info = {}) const&& -> const TensorMatcher&& { + try { + m_verify_impl(view); + } catch (PanicError& e) { + auto oss = std::ostringstream{}; + oss << "Tensor match failed for "; + s_print_tensor(oss, view); + oss << " at " << info.file_name() << ":" << info.line() << "\n- Root cause: " << e.root_cause(); + throw PanicError(std::move(oss).str()); + } + return std::move(*this); + } + + private: + static auto s_print_tensor(std::ostringstream& oss, tvm::ffi::TensorView view) -> void { + oss << "Tensor<"; + int64_t dim = 0; + for (const auto& size : view.shape()) { + if (dim++ > 0) oss << ", "; + oss << size; + } + oss << ">[strides=<"; + dim = 0; + for (const auto& stride : view.strides()) { + if (dim++ > 0) { + oss << ", "; + } + oss << stride; + } + oss << ">, dtype=" << view.dtype(); + oss << ", device=" << details::PrintableDevice{view.device()} << "]"; + } + + auto m_verify_impl(tvm::ffi::TensorView view) const -> void { + const auto dim = static_cast(view.dim()); + RuntimeCheck(dim == m_shape.size(), "Tensor dimension mismatch: expected ", m_shape.size(), " but got ", dim); + for (const auto i : irange(dim)) { + m_shape[i]->verify(view.size(i), "shape", i); + } + if (m_has_strides()) { + for (const auto i : irange(dim)) { + if (view.size(i) != 1 || !m_strides[i]->has_value()) { + // skip stride check for size 1 dimension + m_strides[i]->verify(view.stride(i), "stride", i); + } + } + } else { + RuntimeCheck(view.is_contiguous(), "Tensor is not contiguous as expected"); + } + // since we may double verify, we will force to check + m_dtype->verify(view.dtype()); + m_device->verify(view.device()); + } + + auto m_init_dtype() -> void { + RuntimeCheck(!m_has_dtype, "DType already specified"); + m_has_dtype = true; + } + + auto m_init_device() -> void { + RuntimeCheck(!m_has_device, "Device already specified"); + m_has_device = true; + } + + auto m_has_strides() const -> bool { + return !m_strides.empty(); + } + + std::span m_shape; + std::span m_strides; + DTypeRef m_dtype; + DeviceRef m_device; + bool m_has_dtype = false; + bool m_has_device = false; +}; + +} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh new file mode 100644 index 0000000000..1adc821706 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh @@ -0,0 +1,62 @@ +/// \file tile.cuh +/// \brief Tiled memory access helpers for coalesced global memory I/O. +/// +/// `tile::Memory` represents a contiguous memory region where multiple +/// threads cooperatively load/store elements. The three factory methods +/// determine the thread group: +/// - `thread()` - single thread (no tiling). +/// - `warp()` - all threads in a warp cooperate. +/// - `cta()` - all threads in the CTA cooperate. + +#pragma once +#include + +#include + +namespace device::tile { + +/** + * \brief Represents a contiguous memory region for cooperative tiled access. + * + * Each instance is parameterized by an element type `T` and bound to a + * specific thread id (`tid`) within a group of `tsize` threads. + * + * \tparam T The storage element type (e.g. `AlignedVector, 4>`). + */ +template +struct Memory { + public: + SGL_DEVICE constexpr Memory(uint32_t tid, uint32_t tsize) : tid(tid), tsize(tsize) {} + /// \brief Create a Memory accessor for a single thread (no cooperation). + SGL_DEVICE static constexpr Memory thread() { + return Memory{0, 1}; + } + /// \brief Create a Memory accessor distributed across warp threads. + SGL_DEVICE static Memory warp(int warp_threads = kWarpThreads) { + return Memory{static_cast(threadIdx.x % warp_threads), static_cast(warp_threads)}; + } + /// \brief Create a Memory accessor distributed across all CTA threads. + SGL_DEVICE static Memory cta(int cta_threads = blockDim.x) { + return Memory{static_cast(threadIdx.x), static_cast(cta_threads)}; + } + /// \brief Load one element from `ptr` at the position assigned to this thread. + /// \param ptr Base pointer (cast to `const T*`). + /// \param offset Optional tile offset (multiplied by `tsize`). + SGL_DEVICE T load(const void* ptr, int64_t offset = 0) const { + return static_cast(ptr)[tid + offset * tsize]; + } + /// \brief Store one element to `ptr` at the position assigned to this thread. + SGL_DEVICE void store(void* ptr, T val, int64_t offset = 0) const { + static_cast(ptr)[tid + offset * tsize] = val; + } + /// \brief Check whether this thread's element index is within bounds. + SGL_DEVICE bool in_bound(int64_t element_count, int64_t offset = 0) const { + return tid + offset * tsize < element_count; + } + + private: + uint32_t tid; + uint32_t tsize; +}; + +} // namespace device::tile diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh new file mode 100644 index 0000000000..a7a5346196 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh @@ -0,0 +1,120 @@ +/// \file type.cuh +/// \brief Dtype trait system for CUDA scalar/packed types. +/// +/// `dtype_trait` provides per-type metadata: packed type alias, +/// conversion functions (`from`), and unary/binary math operations. +/// Use `device::cast(from_value)` for type conversion on device. +/// +/// Registered types: +/// | Scalar | Packed (x2) | Notes | +/// |-----------|-------------|-------------------------------| +/// | `fp32_t` | `fp32x2_t` | Full math ops (abs,sqrt,...) | +/// | `fp16_t` | `fp16x2_t` | Conversion only | +/// | `bf16_t` | `bf16x2_t` | Conversion only | +/// | `fp32x2_t`| `fp32x4_t` | Packed float2 <-> half2/bf162 | + +#pragma once +#include + +template +struct dtype_trait {}; + +#define SGL_REGISTER_DTYPE_TRAIT(TYPE, PACK2, ...) \ + template <> \ + struct dtype_trait { \ + using self_t = TYPE; \ + using packed_t = PACK2; \ + template \ + SGL_DEVICE static self_t from(const S& value) { \ + return static_cast(value); \ + } \ + __VA_ARGS__ \ + } + +#define SGL_REGISTER_TYPE_END static_assert(true) + +#define SGL_REGISTER_FROM_FUNCTION(FROM, FN) \ + SGL_DEVICE static self_t from(const FROM& x) { \ + return FN(x); \ + } \ + static_assert(true) + +#define SGL_REGISTER_UNARY_FUNCTION(NAME, FN) \ + SGL_DEVICE static self_t NAME(const self_t& x) { \ + return FN(x); \ + } \ + static_assert(true) + +#define SGL_REGISTER_BINARY_FUNCTION(NAME, FN) \ + SGL_DEVICE static self_t NAME(const self_t& x, const self_t& y) { \ + return FN(x, y); \ + } \ + static_assert(true) + +SGL_REGISTER_DTYPE_TRAIT( + fp32_t, fp32x2_t, SGL_REGISTER_TYPE_END; // + SGL_REGISTER_FROM_FUNCTION(fp16_t, __half2float); + SGL_REGISTER_FROM_FUNCTION(bf16_t, __bfloat162float); + SGL_REGISTER_UNARY_FUNCTION(abs, fabsf); + SGL_REGISTER_UNARY_FUNCTION(sqrt, sqrtf); + SGL_REGISTER_UNARY_FUNCTION(rsqrt, rsqrtf); + SGL_REGISTER_UNARY_FUNCTION(exp, expf); + SGL_REGISTER_UNARY_FUNCTION(sin, sinf); + SGL_REGISTER_UNARY_FUNCTION(cos, cosf); + SGL_REGISTER_BINARY_FUNCTION(max, fmaxf); + SGL_REGISTER_BINARY_FUNCTION(min, fminf);); +SGL_REGISTER_DTYPE_TRAIT(fp16_t, fp16x2_t); +SGL_REGISTER_DTYPE_TRAIT(bf16_t, bf16x2_t); + +/// TODO: Add ROCM implementation +SGL_REGISTER_DTYPE_TRAIT( + fp32x2_t, fp32x4_t, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp16x2_t, __half22float2); + SGL_REGISTER_FROM_FUNCTION(bf16x2_t, __bfloat1622float2);); + +SGL_REGISTER_DTYPE_TRAIT( + fp16x2_t, void, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp32x2_t, __float22half2_rn);); + +SGL_REGISTER_DTYPE_TRAIT( + bf16x2_t, void, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp32x2_t, __float22bfloat162_rn);); + +#ifndef USE_ROCM +SGL_REGISTER_DTYPE_TRAIT(fp8_e4m3_t, fp8x2_e4m3_t); +#endif + +#undef SGL_REGISTER_DTYPE_TRAIT +#undef SGL_REGISTER_FROM_FUNCTION + +/// \brief Alias: the packed (x2) type for `T`. +template +using packed_t = typename dtype_trait::packed_t; + +namespace device { + +/** + * \brief Cast a value from type `From` to type `To` on device. + * + * Dispatches through `dtype_trait::from()`, which uses the appropriate + * CUDA intrinsic (e.g. `__half2float`, `__float22half2_rn`). + */ +template +SGL_DEVICE To cast(const From& value) { + return dtype_trait::from(value); +} + +} // namespace device + +// --------------------------------------------------------------------------- +// FP8 max clamp value — platform-dependent +// CUDA (e4m3fn): 448.0f +// AMD FNUZ (e4m3fnuz): 224.0f +// AMD E4M3 (e4m3fn): 448.0f +// --------------------------------------------------------------------------- +#ifndef USE_ROCM +constexpr float kFP8E4M3Max = 448.0f; +#else // USE_ROCM +#if HIP_FP8_TYPE_FNUZ +constexpr float kFP8E4M3Max = 224.0f; +#else // HIP_FP8_TYPE_E4M3 +constexpr float kFP8E4M3Max = 448.0f; +#endif // HIP_FP8_TYPE_FNUZ +#endif // USE_ROCM diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh new file mode 100644 index 0000000000..2dd6f3dc93 --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh @@ -0,0 +1,333 @@ +/// \file utils.cuh +/// \brief Core CUDA/device utilities: type aliases, PDL helpers, +/// typed pointer access, kernel launch wrapper, and error checking. +/// +/// This header is included (directly or transitively) by nearly every +/// JIT kernel. It provides: +/// - Scalar/packed type aliases (`fp16_t`, `bf16_t`, `fp8_e4m3_t`, ...). +/// - `SGL_DEVICE` macro (forced-inline device function qualifier). +/// - `kWarpThreads` constant (32). +/// - PDL (Programmatic Dependent Launch) helpers for Hopper (sm_90+). +/// - Typed `load_as` / `store_as` for void-pointer access. +/// - `pointer::offset` for safe void-pointer arithmetic. +/// - `host::LaunchKernel` - kernel launcher with optional PDL. +/// - `host::RuntimeDeviceCheck` - CUDA error checking. + +#pragma once + +#include + +#include +#include + +#include +#include +#include +#ifndef USE_ROCM +#include +#include +#include +#include +#else +#include +#include +#include +#ifndef __grid_constant__ +#define __grid_constant__ +#endif +using cudaError_t = hipError_t; +using cudaStream_t = hipStream_t; +using cudaLaunchConfig_t = hipLaunchConfig_t; +using cudaLaunchAttribute = hipLaunchAttribute; +inline constexpr auto cudaSuccess = hipSuccess; +#define cudaStreamPerThread hipStreamPerThread +#define cudaGetErrorString hipGetErrorString +#define cudaGetLastError hipGetLastError +#define cudaLaunchKernel hipLaunchKernel +#define cudaMemcpyAsync hipMemcpyAsync +#define cudaMemcpyHostToDevice hipMemcpyHostToDevice +#define cudaMemcpyDeviceToHost hipMemcpyDeviceToHost +#endif + +#ifndef USE_ROCM +using fp32_t = float; +using fp16_t = __half; +using bf16_t = __nv_bfloat16; +using fp8_e4m3_t = __nv_fp8_e4m3; +using fp8_e5m2_t = __nv_fp8_e5m2; + +using fp32x2_t = float2; +using fp16x2_t = __half2; +using bf16x2_t = __nv_bfloat162; +using fp8x2_e4m3_t = __nv_fp8x2_e4m3; +using fp8x2_e5m2_t = __nv_fp8x2_e5m2; + +using fp32x4_t = float4; +#else +using fp32_t = float; +using fp16_t = __half; +using bf16_t = __hip_bfloat16; +using fp8_e4m3_t = uint8_t; +using fp8_e5m2_t = uint8_t; +using fp32x2_t = float2; +using fp16x2_t = half2; +using bf16x2_t = __hip_bfloat162; +using fp8x2_e4m3_t = uint16_t; +using fp8x2_e5m2_t = uint16_t; +using fp32x4_t = float4; +#endif + +/* + * LDG Support + */ +#ifndef USE_ROCM +#define SGLANG_LDG(arg) __ldg(arg) +#else +#define SGLANG_LDG(arg) *(arg) +#endif + +// DLPack device type for the current platform +#ifndef USE_ROCM +inline constexpr auto kDLGPU = kDLCUDA; +#else +inline constexpr auto kDLGPU = kDLROCM; +#endif + +namespace device { + +/// \brief Macro: forced-inline device function qualifier. +#define SGL_DEVICE __forceinline__ __device__ + +// Architecture detection: SGL_CUDA_ARCH is injected by load_jit() and is +// available in both host and device compilation passes, whereas __CUDA_ARCH__ +// is only defined by nvcc during the device pass. +#if !defined(USE_ROCM) +#if !defined(SGL_CUDA_ARCH) +#error "SGL_CUDA_ARCH is not defined. JIT compilation must inject -DSGL_CUDA_ARCH via load_jit()." +#endif +#if defined(__CUDA_ARCH__) +static_assert( + __CUDA_ARCH__ == SGL_CUDA_ARCH, "SGL_CUDA_ARCH mismatch: injected arch flag does not match device target"); +#endif +#define SGL_ARCH_HOPPER_OR_GREATER (SGL_CUDA_ARCH >= 900) +#define SGL_ARCH_BLACKWELL_OR_GREATER ((SGL_CUDA_ARCH >= 1000) && (CUDA_VERSION >= 12090)) +#else // USE_ROCM +#define SGL_ARCH_HOPPER_OR_GREATER 0 +#define SGL_ARCH_BLACKWELL_OR_GREATER 0 +#endif + +// Maximum vector size in bytes supported by current architecture. +// Pre-Blackwell / AMD: 128-bit (16 bytes) +// Blackwell or greater: 256-bit (32 bytes) +inline constexpr std::size_t kMaxVecBytes = SGL_ARCH_BLACKWELL_OR_GREATER ? 32 : 16; + +/// \brief Number of threads per warp (always 32 on NVIDIA/AMD GPUs). +inline constexpr auto kWarpThreads = 32u; +/// \brief Full warp active mask (all 32 lanes). +#ifndef USE_ROCM +inline constexpr auto kFullMask = 0xffffffffu; +#else +inline constexpr auto kFullMask = 0xffffffffffffffffULL; +#endif + +/** + * \brief PDL (Programmatic Dependent Launch): wait for the primary kernel. + * + * On Hopper (sm_90+), inserts a `griddepcontrol.wait` instruction to + * synchronize with a preceding kernel in the same stream. On older + * architectures or ROCm this is a no-op. + */ +template +SGL_DEVICE void PDLWaitPrimary() { +#if SGL_ARCH_HOPPER_OR_GREATER + if constexpr (kUsePDL) { + asm volatile("griddepcontrol.wait;" ::: "memory"); + } +#endif +} + +/** + * \brief PDL: trigger dependent (secondary) kernel launch. + * + * On Hopper (sm_90+), inserts a `griddepcontrol.launch_dependents` + * instruction. On older architectures or ROCm this is a no-op. + */ +template +SGL_DEVICE void PDLTriggerSecondary() { +#if SGL_ARCH_HOPPER_OR_GREATER + if constexpr (kUsePDL) { + asm volatile("griddepcontrol.launch_dependents;" :::); + } +#endif +} + +template +SGL_DEVICE constexpr auto div_ceil(T a, U b) { + return (a + b - 1) / b; +} + +/** + * \brief Load data with the specified type and offset from a void pointer. + * \tparam T The type to load. + * \param ptr The base pointer. + * \param offset The offset in number of elements of type T. + */ +template +SGL_DEVICE T load_as(const void* ptr, int64_t offset = 0) { + return static_cast(ptr)[offset]; +} + +/** + * \brief Store data with the specified type and offset to a void pointer. + * \tparam T The type to store. + * \param ptr The base pointer. + * \param val The value to store. + * \param offset The offset in number of elements of type T. + * \note we use type_identity_t to force the caller to explicitly specify + * the template parameter `T`, which can avoid accidentally using the wrong type. + */ +template +SGL_DEVICE void store_as(void* ptr, std::type_identity_t val, int64_t offset = 0) { + static_cast(ptr)[offset] = val; +} + +/// \brief Safe void-pointer arithmetic (byte-level by default). +namespace pointer { + +// we only allow void * pointer arithmetic for safety + +template +SGL_DEVICE auto offset(void* ptr, U... offset) -> void* { + return static_cast(ptr) + (... + offset); +} + +template +SGL_DEVICE auto offset(const void* ptr, U... offset) -> const void* { + return static_cast(ptr) + (... + offset); +} + +} // namespace pointer + +} // namespace device + +namespace host { + +/** + * \brief Check the CUDA error code and panic with location info on failure. + */ +inline void RuntimeDeviceCheck(::cudaError_t error, DebugInfo location = {}) { + if (error != ::cudaSuccess) { + [[unlikely]]; + ::host::panic(location, "CUDA error: ", ::cudaGetErrorString(error)); + } +} + +/// \brief Check the last CUDA error (calls `cudaGetLastError`). +inline void RuntimeDeviceCheck(DebugInfo location = {}) { + return RuntimeDeviceCheck(::cudaGetLastError(), location); +} + +/** + * \brief Kernel launcher with automatic stream resolution and PDL support. + * + * Usage: + * \code + * host::LaunchKernel(grid, block, device) + * .enable_pdl(true) + * (my_kernel, arg1, arg2); + * \endcode + * + * The constructor resolves the CUDA stream from a `DLDevice` (via + * `TVMFFIEnvGetStream`) or accepts a raw `cudaStream_t`. The call + * operator launches the kernel and checks for errors. + */ +struct LaunchKernel { + public: + explicit LaunchKernel( + dim3 grid_dim, + dim3 block_dim, + DLDevice device, + std::size_t dynamic_shared_mem_bytes = 0, + DebugInfo location = {}) noexcept + : m_config(s_make_config(grid_dim, block_dim, resolve_device(device), dynamic_shared_mem_bytes)), + m_location(location) {} + + explicit LaunchKernel( + dim3 grid_dim, + dim3 block_dim, + cudaStream_t stream, + std::size_t dynamic_shared_mem_bytes = 0, + DebugInfo location = {}) noexcept + : m_config(s_make_config(grid_dim, block_dim, stream, dynamic_shared_mem_bytes)), m_location(location) {} + + LaunchKernel(const LaunchKernel&) = delete; + LaunchKernel& operator=(const LaunchKernel&) = delete; + + static auto resolve_device(DLDevice device) -> cudaStream_t { + return static_cast(::TVMFFIEnvGetStream(device.device_type, device.device_id)); + } + + auto enable_pdl(bool enabled = true) -> LaunchKernel& { +#ifdef USE_ROCM + (void)enabled; + m_config.numAttrs = 0; +#else + if (enabled) { + auto& attr = m_attrs[m_config.numAttrs++]; + attr.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attr.val.programmaticStreamSerializationAllowed = true; + m_config.attrs = m_attrs; + } +#endif + return *this; + } + + auto enable_cluster(dim3 cluster_dim) -> LaunchKernel& { +#ifdef USE_ROCM + (void)cluster_dim; +#else + auto& attr = m_attrs[m_config.numAttrs++]; + attr.id = cudaLaunchAttributeClusterDimension; + attr.val.clusterDim = {cluster_dim.x, cluster_dim.y, cluster_dim.z}; + m_config.attrs = m_attrs; +#endif + return *this; + } + + template + auto operator()(T&& kernel, Args&&... args) const -> void { +#ifdef USE_ROCM + hipLaunchKernelGGL( + std::forward(kernel), + m_config.gridDim, + m_config.blockDim, + m_config.dynamicSmemBytes, + m_config.stream, + std::forward(args)...); + RuntimeDeviceCheck(m_location); +#else + RuntimeDeviceCheck(::cudaLaunchKernelEx(&m_config, kernel, std::forward(args)...), m_location); +#endif + } + + private: + static auto s_make_config( // Make a config for kernel launch + dim3 grid_dim, + dim3 block_dim, + cudaStream_t stream, + std::size_t smem) -> cudaLaunchConfig_t { + auto config = ::cudaLaunchConfig_t{}; + config.gridDim = grid_dim; + config.blockDim = block_dim; + config.dynamicSmemBytes = smem; + config.stream = stream; + config.numAttrs = 0; + return config; + } + + cudaLaunchConfig_t m_config; + const DebugInfo m_location; + cudaLaunchAttribute m_attrs[2]; +}; + +} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h new file mode 100644 index 0000000000..3226f79ddc --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h @@ -0,0 +1,186 @@ +/// \file utils.h +/// \brief Host-side C++ utilities used by JIT kernel wrappers. +/// +/// Provides: +/// - `DebugInfo` - wraps `std::source_location` for error reporting. +/// - `RuntimeCheck` - runtime assertion with formatted error messages. +/// - `Panic` - unconditional abort with formatted error messages. +/// - `pointer::offset` - safe void-pointer arithmetic (host side). +/// - `div_ceil` - integer ceiling division. +/// - `dtype_bytes` - byte width of a `DLDataType`. +/// - `irange` - Python-style integer range for range-for loops. + +#pragma once + +// ref: https://forums.developer.nvidia.com/t/c-20s-source-location-compilation-error-when-using-nvcc-12-1/258026/3 +#ifdef __CUDACC__ +#include +#if CUDA_VERSION <= 12010 + +#pragma push_macro("__cpp_consteval") +#pragma push_macro("_NODISCARD") +#pragma push_macro("__builtin_LINE") + +#pragma clang diagnostic push +#pragma clang diagnostic ignored "-Wbuiltin-macro-redefined" +#define __cpp_consteval 201811L +#pragma clang diagnostic pop + +#ifdef _NODISCARD +#undef _NODISCARD +#define _NODISCARD +#endif + +#define consteval constexpr + +#include "source_location.h" + +#undef consteval +#pragma pop_macro("__cpp_consteval") +#pragma pop_macro("_NODISCARD") +#else // __CUDACC__ && CUDA_VERSION > 12010 +#include "source_location.h" +#endif +#else // no __CUDACC__ +#include "source_location.h" +#endif + +#include + +#include +#include +#include +#include +#include +#include + +namespace host { + +template +inline constexpr bool dependent_false_v = false; + +/// \brief Source-location wrapper for debug/error messages. +struct DebugInfo : public source_location_t { + DebugInfo(source_location_t loc = source_location_t::current()) : source_location_t(loc) {} +}; + +/// \brief Exception type thrown by `RuntimeCheck` and `Panic`. +struct PanicError : public std::runtime_error { + public: + explicit PanicError(std::string msg) : runtime_error(msg), m_message(std::move(msg)) {} + auto root_cause() const -> std::string_view { + const auto str = std::string_view{m_message}; + const auto pos = str.find(": "); + return pos == std::string_view::npos ? str : str.substr(pos + 2); + } + + private: + std::string m_message; +}; + +/// \brief Unconditionally abort with a formatted error message. +template +[[noreturn]] +inline auto panic(DebugInfo location, Args&&... args) -> void { + std::ostringstream os; + os << "Runtime check failed at " << location.file_name() << ":" << location.line(); + if constexpr (sizeof...(args) > 0) { + os << ": "; + (os << ... << std::forward(args)); + } else { + os << " in " << location.function_name(); + } + throw PanicError(std::move(os).str()); +} + +/** + * \brief Runtime assertion: panics with a formatted message when `condition` + * is false. Extra `args` are streamed to the error message. + * + * Example: + * \code + * RuntimeCheck(n > 0, "n must be positive, got ", n); + * \endcode + */ +template +struct RuntimeCheck { + template + explicit RuntimeCheck(Cond&& condition, Args&&... args, DebugInfo location = {}) { + if (condition) return; + [[unlikely]] ::host::panic(location, std::forward(args)...); + } + template + explicit RuntimeCheck(DebugInfo location, Cond&& condition, Args&&... args) { + if (condition) return; + [[unlikely]] ::host::panic(location, std::forward(args)...); + } +}; + +template +struct Panic { + explicit Panic(Args&&... args, DebugInfo location = {}) { + ::host::panic(location, std::forward(args)...); + } + explicit Panic(DebugInfo location, Args&&... args) { + ::host::panic(location, std::forward(args)...); + } + [[noreturn]] ~Panic() { + std::terminate(); + } +}; + +template +explicit RuntimeCheck(Cond&&, Args&&...) -> RuntimeCheck; + +template +explicit RuntimeCheck(DebugInfo, Cond&&, Args&&...) -> RuntimeCheck; + +template +explicit Panic(Args&&...) -> Panic; + +template +explicit Panic(DebugInfo, Args&&...) -> Panic; + +namespace pointer { + +// we only allow void * pointer arithmetic for safety + +template +inline auto offset(void* ptr, U... offset) -> void* { + return static_cast(ptr) + (... + offset); +} + +template +inline auto offset(const void* ptr, U... offset) -> const void* { + return static_cast(ptr) + (... + offset); +} + +} // namespace pointer + +/// \brief Integer ceiling division: ceil(a / b). +template +inline constexpr auto div_ceil(T a, U b) { + return (a + b - 1) / b; +} + +/// \brief Returns the byte width of a DLPack data type. +inline auto dtype_bytes(DLDataType dtype) -> std::size_t { + return static_cast(dtype.bits / 8); +} + +namespace stdr = std::ranges; +namespace stdv = stdr::views; + +/// \brief Python-style integer range: `irange(n)` -> `[0, n)`. +template +inline auto irange(T end) { + return stdv::iota(static_cast(0), end); +} + +/// \brief Python-style integer range: `irange(start, end)` -> `[start, end)`. +template +inline auto irange(T start, T end) { + return stdv::iota(start, end); +} + +} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh new file mode 100644 index 0000000000..67f388679f --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh @@ -0,0 +1,118 @@ +/// \file vec.cuh +/// \brief Aligned vector types for coalesced global memory access. +/// +/// `AlignedVector` wraps `N` elements of type `T` in a naturally +/// aligned struct so that the compiler emits wide (vectorized) load/store +/// instructions (e.g. `LDG.128`). The maximum supported vector width is +/// 256 bits (32 bytes), matching CUDA's widest vector load. + +#pragma once +#include + +#include +#include + +namespace device { + +namespace details { + +/// \brief Maps byte-width to the corresponding unsigned integer type. +template +struct uint_trait {}; + +template <> +struct uint_trait<1> { + using type = uint8_t; +}; + +template <> +struct uint_trait<2> { + using type = uint16_t; +}; + +template <> +struct uint_trait<4> { + using type = uint32_t; +}; + +template <> +struct uint_trait<8> { + using type = uint64_t; +}; + +/// \brief Alias: maps `sizeof(T)` to matching unsigned int type. +template +using sized_int = typename uint_trait::type; + +} // namespace details + +/// \brief Raw aligned storage for `N` elements of type `T`. +template +struct alignas(sizeof(T) * N) AlignedStorage { + T data[N]; +}; + +/** + * \brief Aligned vector for vectorized memory access on GPU. + * + * Stores `N` elements of type `T` with natural alignment so that a single + * `load`/`store` call compiles to a wide memory transaction. + * + * \tparam T Element type (e.g. `fp16_t`, `bf16_t`, `float`). + * \tparam N Number of elements. Must be a power of two and + * `sizeof(T) * N <= 32` (256 bits). + * + * Example: + * \code + * AlignedVector vec; // 16 bytes, 128-bit aligned + * vec.load(input_ptr, tid); // vectorized load + * vec[0] = vec[0] + 1; + * vec.store(output_ptr, tid); // vectorized store + * \endcode + */ +template +struct AlignedVector { + private: + static_assert( + (N > 0 && (N & (N - 1)) == 0) && sizeof(T) * N <= kMaxVecBytes, + "CUDA vector size exceeds arch limit: max 16 bytes on pre-Blackwell/AMD, " + "32 bytes on Blackwell or greater"); + using element_t = typename details::sized_int; + using storage_t = AlignedStorage; + + public: + /// \brief Vectorized load from `ptr` at the given element `offset`. + SGL_DEVICE void load(const void* ptr, int64_t offset = 0) { + m_storage = reinterpret_cast(ptr)[offset]; + } + /// \brief Vectorized store to `ptr` at the given element `offset`. + SGL_DEVICE void store(void* ptr, int64_t offset = 0) const { + reinterpret_cast(ptr)[offset] = m_storage; + } + /// \brief Fill all N elements with the same `value`. + SGL_DEVICE void fill(T value) { + const auto store_value = *reinterpret_cast(&value); +#pragma unroll + for (std::size_t i = 0; i < N; ++i) { + m_storage.data[i] = store_value; + } + } + + SGL_DEVICE auto operator[](std::size_t idx) -> T& { + return reinterpret_cast(&m_storage)[idx]; + } + SGL_DEVICE auto operator[](std::size_t idx) const -> T { + return reinterpret_cast(&m_storage)[idx]; + } + SGL_DEVICE auto data() -> T* { + return reinterpret_cast(&m_storage); + } + SGL_DEVICE auto data() const -> const T* { + return reinterpret_cast(&m_storage); + } + + private: + storage_t m_storage; +}; + +} // namespace device diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh new file mode 100644 index 0000000000..9d82efae1e --- /dev/null +++ b/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh @@ -0,0 +1,56 @@ +/// \file warp.cuh +/// \brief Warp-level reduction primitives. + +#pragma once +#include +#include + +namespace device::warp { + +/// \brief Full warp active mask. +#ifndef USE_ROCM +static constexpr uint32_t kFullMask = 0xffffffffu; +using mask_t = uint32_t; +#else +static constexpr uint64_t kFullMask = 0xffffffffffffffffULL; +using mask_t = uint64_t; +#endif + +/** + * \brief Warp-level sum reduction. + * + * On CUDA: uses __shfl_xor_sync with width=32. + * On HIP: uses __shfl_xor with explicit width parameter (supports wave64 sub-groups). + */ +template +SGL_DEVICE T reduce_sum(T value, mask_t active_mask = kFullMask) { + static_assert(kNumThreads >= 1 && kNumThreads <= kWarpThreads); + static_assert(std::has_single_bit(kNumThreads), "must be pow of 2"); +#pragma unroll + for (int mask = kNumThreads / 2; mask > 0; mask >>= 1) +#ifndef USE_ROCM + value = value + __shfl_xor_sync(active_mask, value, mask, 32); +#else + value = value + __shfl_xor(value, mask, kNumThreads); +#endif + return value; +} + +/** + * \brief Warp-level max reduction. + */ +template +SGL_DEVICE T reduce_max(T value, mask_t active_mask = kFullMask) { + static_assert(kNumThreads >= 1 && kNumThreads <= kWarpThreads); + static_assert(std::has_single_bit(kNumThreads), "must be pow of 2"); +#pragma unroll + for (int mask = kNumThreads / 2; mask > 0; mask >>= 1) +#ifndef USE_ROCM + value = math::max(value, __shfl_xor_sync(active_mask, value, mask, 32)); +#else + value = math::max(value, __shfl_xor(value, mask, kNumThreads)); +#endif + return value; +} + +} // namespace device::warp diff --git a/lightllm/third_party/sglang_jit/jit_utils.py b/lightllm/third_party/sglang_jit/jit_utils.py new file mode 100644 index 0000000000..4096c16bb4 --- /dev/null +++ b/lightllm/third_party/sglang_jit/jit_utils.py @@ -0,0 +1,432 @@ +from __future__ import annotations + +import functools +import importlib.util +import logging +import os +import pathlib +from contextlib import contextmanager +from dataclasses import dataclass +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Dict, + List, + Optional, + Tuple, + TypeAlias, + TypeVar, + Union, +) + +import torch + +if TYPE_CHECKING: + from tvm_ffi import Module + +F = TypeVar("F", bound=Callable[..., Any]) +_FULL_TEST_ENV_VAR = "SGLANG_JIT_KERNEL_RUN_FULL_TESTS" + +logger = logging.getLogger(__name__) + + +def is_in_ci() -> bool: + return os.getenv("SGLANG_IS_IN_CI", "").lower() in ("1", "true", "yes", "y") + + +def should_run_full_tests() -> bool: + return os.getenv(_FULL_TEST_ENV_VAR, "false").lower() == "true" + + +def get_ci_test_range(full_range: List[Any], ci_range: List[Any]) -> List[Any]: + if should_run_full_tests(): + return full_range + return ci_range if is_in_ci() else full_range + + +def cache_once(fn: F) -> F: + """ + NOTE: `functools.lru_cache` is not compatible with `torch.compile` + So we manually implement a simple cache_once decorator to replace it. + """ + result_map = {} + + @functools.wraps(fn) + def wrapper(*args, **kwargs): + key = (args, tuple(sorted(kwargs.items()))) + if key not in result_map: + result_map[key] = fn(*args, **kwargs) + return result_map[key] + + return wrapper # type: ignore + + +def _make_wrapper(tup: Tuple[str, str]) -> str: + export_name, kernel_name = tup + return f"TVM_FFI_DLL_EXPORT_TYPED_FUNC({export_name}, ({kernel_name}));" + + +@cache_once +def _resolve_kernel_path() -> pathlib.Path: + cur_dir = pathlib.Path(__file__).parent.resolve() + + # first, try this directory structure + def _environment_install(): + candidate = cur_dir.resolve() + if (candidate / "include").exists() and (candidate / "csrc").exists(): + return candidate + return None + + def _package_install(): + # TODO: support find path by package + return None + + path = _environment_install() or _package_install() + if path is None: + raise RuntimeError("Cannot find sglang.jit_kernel path") + return path + + +KERNEL_PATH = _resolve_kernel_path() +DEFAULT_INCLUDE = [str(KERNEL_PATH / "include")] +DEFAULT_CFLAGS = ["-std=c++20", "-O3"] +DEFAULT_LDFLAGS = [] +CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, str, bool, torch.dtype] + + +class CPPArgList(list[str]): + def __str__(self) -> str: + return ", ".join(self) + + +CPP_DTYPE_MAP = { + torch.float: "fp32_t", + torch.float16: "fp16_t", + torch.float8_e4m3fn: "fp8_e4m3_t", + torch.bfloat16: "bf16_t", + torch.int8: "int8_t", + torch.int32: "int32_t", + torch.int64: "int64_t", +} + + +# AMD/ROCm note: +@cache_once +def is_hip_runtime() -> bool: + return bool(torch.version.hip) + + +# MThreads/MUSA note: +@cache_once +def is_musa_runtime() -> bool: + return hasattr(torch.version, "musa") and torch.version.musa is not None + + +def make_cpp_args(*args: CPP_TEMPLATE_TYPE) -> CPPArgList: + def _convert(arg: CPP_TEMPLATE_TYPE) -> str: + if isinstance(arg, bool): + return "true" if arg else "false" + if isinstance(arg, (int, str, float)): + return str(arg) + if isinstance(arg, torch.dtype): + return CPP_DTYPE_MAP[arg] + raise TypeError(f"Unsupported argument type for cpp template: {type(arg)}") + + return CPPArgList(_convert(arg) for arg in args) + + +def load_jit( + *args: str, + cpp_files: List[str] | None = None, + cuda_files: List[str] | None = None, + cpp_wrappers: List[Tuple[str, str]] | None = None, + cuda_wrappers: List[Tuple[str, str]] | None = None, + extra_cflags: List[str] | None = None, + extra_cuda_cflags: List[str] | None = None, + extra_ldflags: List[str] | None = None, + extra_include_paths: List[str] | None = None, + extra_dependencies: List[str] | None = None, + build_directory: str | None = None, + header_only: bool = True, +) -> Module: + """ + Loading a JIT module from C++/CUDA source files. + We define a wrapper as a tuple of (export_name, kernel_name), + where `export_name` is the name used to called from Python, + and `kernel_name` is the name of the kernel class in C++/CUDA source. + + :param args: Unique marker of the JIT module. Must be distinct for different kernels. + :type args: str + :param cpp_files: A list of C++ source files. + :type cpp_files: List[str] | None + :param cuda_files: A list of CUDA source files. + :type cuda_files: List[str] | None + :param cpp_wrappers: A list of C++ wrappers, defining the export name and kernel name. + :type cpp_wrappers: List[Tuple[str, str]] | None + :param cuda_wrappers: A list of CUDA wrappers, defining the export name and kernel name. + :type cuda_wrappers: List[Tuple[str, str]] | None + :param extra_cflags: Extra C++ compiler flags. + :type extra_cflags: List[str] | None + :param extra_cuda_cflags: Extra CUDA compiler flags. + :type extra_cuda_cflags: List[str] | None + :param extra_ldflags: Extra linker flags. + :type extra_ldflags: List[str] | None + :param extra_include_paths: Extra include paths. + :type extra_include_paths: List[str] | None + :param extra_dependencies: Extra dependencies for the JIT module, e.g., cutlass. + :type extra_dependencies: List[str] | None + :param build_directory: The build directory for JIT compilation. + :type build_directory: str | None + :param header_only: Whether the module is header-only. + If true, apply the wrappers to export given class/functions. + Otherwise, we must export from C++/CUDA side. + :return: A just-in-time(JIT) compiled module. + :rtype: Module + """ + + from tvm_ffi.cpp import load, load_inline + + cpp_files = cpp_files or [] + cuda_files = cuda_files or [] + extra_cflags = extra_cflags or [] + extra_cuda_cflags = extra_cuda_cflags or [] + extra_ldflags = extra_ldflags or [] + extra_include_paths = extra_include_paths or [] + + cpp_files = [str((KERNEL_PATH / "csrc" / f).resolve()) for f in cpp_files] + cuda_files = [str((KERNEL_PATH / "csrc" / f).resolve()) for f in cuda_files] + + for dep in set(extra_dependencies or []): + if dep not in _REGISTERED_DEPENDENCIES: + raise ValueError(f"Dependency {dep} is not registered.") + extra_include_paths += _REGISTERED_DEPENDENCIES[dep]() + + module_name = "sgl_kernel_jit_" + "_".join(str(arg) for arg in args) + if header_only: + cpp_wrappers = cpp_wrappers or [] + cuda_wrappers = cuda_wrappers or [] + cpp_sources = [f'#include "{path}"' for path in cpp_files] + cpp_sources += [_make_wrapper(tup) for tup in cpp_wrappers] + + # include cuda files + cuda_sources = [f'#include "{path}"' for path in cuda_files] + cuda_sources += [_make_wrapper(tup) for tup in cuda_wrappers] + with _jit_compile_context(): + return load_inline( + module_name, + cpp_sources=cpp_sources, + cuda_sources=cuda_sources, + extra_cflags=DEFAULT_CFLAGS + extra_cflags, + extra_cuda_cflags=_get_default_target_flags() + extra_cuda_cflags, + extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags, + extra_include_paths=DEFAULT_INCLUDE + extra_include_paths, + build_directory=build_directory, + ) + else: + assert cpp_wrappers is None and cuda_wrappers is None + with _jit_compile_context(): + return load( + module_name, + cpp_files=cpp_files, + cuda_files=cuda_files, + extra_cflags=DEFAULT_CFLAGS + extra_cflags, + extra_cuda_cflags=_get_default_target_flags() + extra_cuda_cflags, + extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags, + extra_include_paths=DEFAULT_INCLUDE + extra_include_paths, + build_directory=build_directory, + ) + + +@dataclass +class ArchInfo: + major: int + minor: int + suffix: str + + @property + def target_name(self) -> str: + return f"{self.major}.{self.minor}{self.suffix}" + + @property + def jit_flag(self) -> str: + return f"-DSGL_CUDA_ARCH={self.major * 100 + self.minor * 10}" + + +@cache_once +def _init_jit_cuda_arch_once(): + global _CUDA_ARCH + try: + device = torch.cuda.current_device() + major, minor = torch.cuda.get_device_capability(device) + except Exception: + logger.warning("Cannot detect CUDA architecture.") + major, minor = 0, 0 # invalid value to trigger compile error if used + _CUDA_ARCH = ArchInfo(major, minor, "") + + +@contextmanager +def _jit_compile_context(): + if is_hip_runtime(): + yield # TODO: support ROCm `TVM_FFI_ROCM_ARCH_LIST` if needed + return + env_key = "TVM_FFI_CUDA_ARCH_LIST" + old_value = os.environ.get(env_key, None) + os.environ[env_key] = get_jit_cuda_arch().target_name + try: + yield + finally: + if old_value is None: + os.environ.pop(env_key, None) + else: + os.environ[env_key] = old_value + + +# NOTE: this might also be used in __main__.py for compile flags export +def _get_default_target_flags() -> List[str]: + if is_hip_runtime(): + flags = ["-DUSE_ROCM", "-std=c++20", "-O3"] + # Detect FP8 type based on GPU architecture + try: + device = torch.cuda.current_device() + gcn_arch = torch.cuda.get_device_properties(device).gcnArchName + if "gfx942" in gcn_arch: + flags.append("-DHIP_FP8_TYPE_FNUZ=1") + else: + flags.append("-DHIP_FP8_TYPE_E4M3=1") + except Exception: + flags.append("-DHIP_FP8_TYPE_E4M3=1") + return flags + else: + return [ + get_jit_cuda_arch().jit_flag, + "-std=c++20", + "-O3", + "--expt-relaxed-constexpr", + ] + + +@contextmanager +def override_jit_cuda_arch(major: int, minor: int, suffix: str = ""): + """A context manager to temporarily override CUDA architecture.""" + global _CUDA_ARCH + old_value = get_jit_cuda_arch() + _CUDA_ARCH = ArchInfo(major, minor, suffix) + try: + yield + finally: + _CUDA_ARCH = old_value + + +def get_jit_cuda_arch() -> ArchInfo: + """Get the current CUDA architecture info.""" + _init_jit_cuda_arch_once() + return _CUDA_ARCH + + +@cache_once +def is_arch_support_pdl() -> bool: + if is_hip_runtime() or is_musa_runtime(): + return False + return get_jit_cuda_arch().major >= 9 + + +def _find_package_root(package: str) -> Optional[pathlib.Path]: + spec = importlib.util.find_spec(package) + if spec is None or spec.origin is None: + return None + return pathlib.Path(spec.origin).resolve().parent + + +# NOTE: this might also be used in __main__.py for compile flags export +_REGISTERED_DEPENDENCIES: Dict[str, Callable[[], List[str]]] = {} + + +def register_dependency(name: str): + def decorator(f: Callable[[], List[str]]) -> Callable[[], List[str]]: + if name in _REGISTERED_DEPENDENCIES: + raise ValueError(f"Dependency {name} already registered") + _REGISTERED_DEPENDENCIES[name] = f + return f + + return decorator + + +@register_dependency("flashinfer") +def get_flashinfer_include_paths() -> List[str]: + include_paths: List[str] = [] + flashinfer_root = _find_package_root("flashinfer") + if flashinfer_root is None: + raise RuntimeError( + "Cannot find flashinfer package. Please install flashinfer to get" + "the required headers for JIT compilation." + ) + + flashinfer_data = flashinfer_root / "data" + candidates = [ + flashinfer_data / "include", + flashinfer_data / "csrc", + flashinfer_data / "cutlass" / "include", + flashinfer_data / "cutlass" / "tools" / "util" / "include", + flashinfer_data / "spdlog" / "include", + ] + + for path in candidates: + if not path.exists(): + raise RuntimeError( + f"Required header path {path} for flashinfer dependency not found." + " Please check your flashinfer installation." + ) + include_paths.append(str(path)) + return include_paths + + +@register_dependency("cutlass") +def get_cutlass_include_paths() -> List[str]: + include_paths: List[str] = [] + + flashinfer_root = _find_package_root("flashinfer") + if flashinfer_root is not None: + candidates = [ + flashinfer_root / "data" / "cutlass" / "include", + flashinfer_root / "data" / "cutlass" / "tools" / "util" / "include", + ] + for path in candidates: + if path.exists(): + include_paths.append(str(path)) + + deep_gemm_root = _find_package_root("deep_gemm") + if deep_gemm_root is not None: + candidate = deep_gemm_root / "include" + if candidate.exists(): + include_paths.append(str(candidate)) + + # De-duplicate while preserving order. + unique_paths = [] + seen = set() + for path in include_paths: + if path in seen: + continue + seen.add(path) + unique_paths.append(path) + + if not unique_paths: + raise RuntimeError( + "Cannot find CUTLASS headers required for JIT compilation. " + "Please install flashinfer or deep_gemm with CUTLASS headers." + ) + return unique_paths + + +__all__ = [ + "should_run_full_tests", + "get_ci_test_range", + "cache_once", + "is_hip_runtime", + "make_cpp_args", + "load_jit", + "override_jit_cuda_arch", + "get_jit_cuda_arch", + "is_arch_support_pdl", + "register_dependency", +] diff --git a/lightllm/third_party/sglang_jit/runtime_utils.py b/lightllm/third_party/sglang_jit/runtime_utils.py new file mode 100644 index 0000000000..d322498ca4 --- /dev/null +++ b/lightllm/third_party/sglang_jit/runtime_utils.py @@ -0,0 +1,5 @@ +import torch + + +def is_hip() -> bool: + return torch.version.hip is not None From e8c49d101b31ffe07f67ae9c5738f3e42a0ca810 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 18 Jun 2026 05:44:58 +0000 Subject: [PATCH 030/214] fix tpsp --- .../layer_infer/transformer_layer_infer.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 617d0dcd85..57b171a8d1 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -166,7 +166,7 @@ def _get_qkv( freqs_cis=self.freqs_cis, positions=infer_state.position_ids, ) - return q, qa + return q, qa, input def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): # o: [T, tp_q_head_num_, head_dim_] after inverse rope -> grouped low-rank O -> [T, embed_dim_] @@ -185,8 +185,8 @@ def context_attention_forward( # _get_qkv writes the chunk's packed latent into the swa pool (fused kernel) before # attention reads it back via full_to_swa indices (this custom forward bypasses the # tpl _post_cache_kv path). - q, q_lora = self._get_qkv(x, infer_state, layer_weight) - o = self._context_attention_wrapper_run(q, q_lora, x, infer_state, layer_weight) + q, q_lora, full_x = self._get_qkv(x, infer_state, layer_weight) + o = self._context_attention_wrapper_run(q, q_lora, full_x, infer_state, layer_weight) return self._get_o(o, infer_state, layer_weight) def _context_attention_wrapper_run( @@ -262,8 +262,8 @@ def _context_attention_kernel( def token_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - q, q_lora = self._get_qkv(x, infer_state, layer_weight) - o = self._token_attention_kernel(q, q_lora, x, infer_state, layer_weight) + q, q_lora, full_x = self._get_qkv(x, infer_state, layer_weight) + o = self._token_attention_kernel(q, q_lora, full_x, infer_state, layer_weight) return self._get_o(o, infer_state, layer_weight) def _token_attention_kernel( @@ -552,7 +552,12 @@ def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, la cos_tok = infer_state.position_cos_compress sin_tok = infer_state.position_sin_compress - idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(x.shape[0], self.index_n_heads, self.index_head_dim) + token_num = q_lora.shape[0] + if x.shape[0] != token_num: + raise RuntimeError( + f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" + ) + idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(token_num, self.index_n_heads, self.index_head_dim) rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) idx_q = hadamard_transform(idx_q, scale=self.index_head_dim ** -0.5) idx_q_fp8, q_scale = act_quant(idx_q, self.index_head_dim, None) # fp8 [T,H,d], scale [T,H,1] From 88309b5ffd71754683945596d5a768a603fc3161 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 21 Jun 2026 23:57:11 +0000 Subject: [PATCH 031/214] fix profile --- .../deepseek4_mem_manager.py | 173 +++++------------- lightllm/common/req_manager.py | 22 --- lightllm/models/deepseek_v4/model.py | 7 +- 3 files changed, 47 insertions(+), 155 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index cfb149dcec..8dedf16387 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,13 +1,11 @@ import torch -import torch.distributed as dist from typing import List, Optional, Union from .mem_manager import MemoryManager from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator -from lightllm.utils.dist_utils import get_current_device_id, get_current_rank_in_node +from lightllm.utils.dist_utils import get_current_rank_in_node from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger -from lightllm.utils.profile_max_tokens import get_available_gpu_memory, get_total_gpu_memory logger = init_logger(__name__) @@ -15,30 +13,29 @@ # fp8_ds_mla packed-latent byte layout (ABI shared with the flash_mla extra-cache fork and # sglang/vllm): 448B NoPE fp8 + 64*2B RoPE bf16 + 7B ue8m0 scale + 1B pad = 584B per token, # stored in page slabs whose tail carries the per-token scale bytes. -DSV4_MLA_NOPE_DIM = 448 -DSV4_MLA_ROPE_DIM = 64 -DSV4_MLA_HEAD_DIM = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM -DSV4_MLA_QUANT_GROUP_SIZE = 64 -DSV4_MLA_SCALE_BYTES = DSV4_MLA_NOPE_DIM // DSV4_MLA_QUANT_GROUP_SIZE + 1 -DSV4_MLA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 + DSV4_MLA_SCALE_BYTES -DSV4_MLA_DATA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 -DSV4_MLA_PAGE_ALIGN_BYTES = DSV4_MLA_DATA_BYTES_PER_TOKEN -DSV4_INDEXER_HEAD_DIM = 128 -DSV4_INDEXER_SCALE_BYTES = 4 -DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + DSV4_INDEXER_SCALE_BYTES -DSV4_FP8_E4M3_MAX = 448.0 -DSV4_FP8_SCALE_MIN = 1e-4 -DSV4_SWA_PAGE_SIZE = 128 -DSV4_C4_PAGE_SIZE = 64 -DSV4_C128_PAGE_SIZE = 2 -DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 +DSV4_MLA_NOPE_DIM = 448 # 448B +DSV4_MLA_ROPE_DIM = 64 # 64 dim +DSV4_MLA_HEAD_DIM = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM # 512 +DSV4_MLA_QUANT_GROUP_SIZE = 64 # 64 +DSV4_MLA_SCALE_BYTES = DSV4_MLA_NOPE_DIM // DSV4_MLA_QUANT_GROUP_SIZE + 1 # 8 (7 ue8m0 + 1 pad) +DSV4_MLA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 + DSV4_MLA_SCALE_BYTES # 584 +DSV4_MLA_DATA_BYTES_PER_TOKEN = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM * 2 # 576 +DSV4_MLA_PAGE_ALIGN_BYTES = DSV4_MLA_DATA_BYTES_PER_TOKEN # 576 +DSV4_INDEXER_HEAD_DIM = 128 # 128 +DSV4_INDEXER_SCALE_BYTES = 4 # 4B fp32 scale +DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + DSV4_INDEXER_SCALE_BYTES # 132 +DSV4_FP8_E4M3_MAX = 448.0 # 448.0 +DSV4_FP8_SCALE_MIN = 1e-4 # 1e-4 +DSV4_SWA_PAGE_SIZE = 128 # 128 slots/page +DSV4_C4_PAGE_SIZE = 64 # 64 slots/page +DSV4_C128_PAGE_SIZE = 2 # 2 slots/page +DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 # 256 (= c4 ratio) # compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 离线聚合为每页 1 组。 -DSV4_C4_STATE_RING = 8 -DSV4_C128_STATE_RING = 128 -# swa 池占 full token 空间的比例下限(sglang swa_full_tokens_ratio=0.1 同值)。 -# v5 的 swa 压力阀(借页/驱逐)已覆盖 radix 树与准入波次的瞬时增长,结构性预算 -# (max_req×window + batch_max_tokens 余量)另行叠加,0.1 仅作 full 池比例下限。 -DSV4_SWA_FULL_TOKENS_RATIO = 0.1 +DSV4_C4_STATE_RING = 8 # 8 rows/page +DSV4_C128_STATE_RING = 128 # 128 rows/page +# swa 池占 full token 空间的比例(sglang DSV4 默认 swa_full_tokens_ratio=0.1 同值)。 +# 瞬时借页/驱逐走 swa 压力阀;池子大小仅按 ratio 切分,不再叠加结构性余量。 +DSV4_SWA_FULL_TOKENS_RATIO = 0.1 # 0.1 def _ceil_div(a: int, b: int) -> int: @@ -134,14 +131,14 @@ class DeepseekV4MemoryManager(MemoryManager): operator_class = DeepseekV4MemOperator - mla_nope_dim = DSV4_MLA_NOPE_DIM - mla_rope_dim = DSV4_MLA_ROPE_DIM - mla_head_dim = DSV4_MLA_HEAD_DIM - mla_quant_group_size = DSV4_MLA_QUANT_GROUP_SIZE - mla_scale_bytes = DSV4_MLA_SCALE_BYTES - mla_bytes_per_token = DSV4_MLA_BYTES_PER_TOKEN - indexer_head_dim_default = DSV4_INDEXER_HEAD_DIM - indexer_bytes_per_token = DSV4_INDEXER_BYTES_PER_TOKEN + mla_nope_dim = DSV4_MLA_NOPE_DIM # 448 + mla_rope_dim = DSV4_MLA_ROPE_DIM # 64 + mla_head_dim = DSV4_MLA_HEAD_DIM # 512 + mla_quant_group_size = DSV4_MLA_QUANT_GROUP_SIZE # 64 + mla_scale_bytes = DSV4_MLA_SCALE_BYTES # 8 + mla_bytes_per_token = DSV4_MLA_BYTES_PER_TOKEN # 584 + indexer_head_dim_default = DSV4_INDEXER_HEAD_DIM # 128 + indexer_bytes_per_token = DSV4_INDEXER_BYTES_PER_TOKEN # 132 def __init__( self, @@ -154,7 +151,6 @@ def __init__( indexer_head_dim: int = 128, max_request_num: Optional[int] = None, sliding_window: Optional[int] = None, - swa_extra_token_num: int = 0, swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, @@ -173,9 +169,6 @@ def __init__( self.indexer_head_dim = indexer_head_dim self.max_request_num = max_request_num self.sliding_window = sliding_window - # 活跃窗口(max_request_num * sliding_window)之外的余量: 在途 prefill chunk 的瞬时占用 - # (出窗槽位要到下一次 prep 才回收) + radix cache 持有的窗口尾部。 - self.swa_extra_token_num = int(swa_extra_token_num) self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) @@ -193,26 +186,12 @@ def __init__( super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) # ------------------------------------------------------------------ sizing - def _swa_per_req_budget(self) -> int: - # 活跃请求保留 window + 一个 radix 页(req_manager._swa_retain_len: 让最近完成的 - # prompt-cache 边界的结尾页恒驻留)。V4 的 prompt-cache 边界取 256 token, - # 避免 radix 共享前缀落在 c4 物理页(64 c4 entry = 256 token)中间。 - return int(self.sliding_window) + DSV4_PROMPT_CACHE_PAGE_SIZE - def _planned_swa_size(self, full_size: int) -> int: - # swa 池按页分配(页 = 128 = sliding_window),容量向上取整到整页。 - if self.max_request_num is None or self.sliding_window is None: - return _ceil_div(full_size, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE - cap = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num - cap = max(cap, int(full_size * self.swa_full_tokens_ratio)) + # swa 池 = full * ratio,按 SWA 物理页(128)向上取整;与 sglang DSV4PoolConfigurator 同思路。 + cap = int(full_size * self.swa_full_tokens_ratio) cap = max(1, min(full_size, cap)) return _ceil_div(cap, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE - @staticmethod - def _slab_bytes_per_slot(page_size: int, data_bytes: int, scale_bytes: int, align_bytes: int = 1) -> float: - bytes_per_page = _ceil_div(page_size * (data_bytes + scale_bytes), align_bytes) * align_bytes - return bytes_per_page / page_size - @staticmethod def _paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: rows = num_swa_pages * ring + ring + 1 @@ -225,79 +204,20 @@ def _init_state_sentinel(buffer: torch.Tensor) -> None: buffer[:, -1, half:].fill_(float("-inf")) return - def _paged_state_bytes_per_swa_slot(self) -> float: - """c4/c128 compressor state(swa 页派生寻址)摊到每个 swa 槽的字节数。""" - per_page = 0.0 - if self.n_c4 > 0: - per_page += DSV4_C4_STATE_RING * (4 * self.head_dim + 4 * self.indexer_head_dim) * 4 * self.n_c4 - if self.n_c128 > 0: - per_page += DSV4_C128_STATE_RING * (2 * self.head_dim) * 4 * self.n_c128 - return per_page / DSV4_SWA_PAGE_SIZE - - def _swa_slot_bytes(self) -> float: - per_layer = self._slab_bytes_per_slot( - DSV4_SWA_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES - ) - return per_layer * self.layer_num + self._paged_state_bytes_per_swa_slot() - - def _compressed_cell_size(self) -> float: - """每个 full token 摊到压缩池上的精确字节数(按 page-slab 对齐后)。""" - c4_latent = self._slab_bytes_per_slot( - DSV4_C4_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES - ) - c128_latent = self._slab_bytes_per_slot( - DSV4_C128_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES - ) - c4_indexer = self._slab_bytes_per_slot(DSV4_C4_PAGE_SIZE, self.indexer_head_dim, DSV4_INDEXER_SCALE_BYTES) - return (c4_latent + c4_indexer) * self.n_c4 / 4 + c128_latent * self.n_c128 / 128 - def get_cell_size(self): - compressed = self._compressed_cell_size() - if self.size is None: - return self._swa_slot_bytes() + compressed - swa_ratio = self._planned_swa_size(self.size) / max(1, self.size) - return self._swa_slot_bytes() * swa_ratio + compressed - - def profile_size(self, mem_fraction): - if self.size is not None: - return - - torch.cuda.empty_cache() - world_size = dist.get_world_size() - available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) - available_bytes = available_memory * 1024 ** 3 - swa_slot_bytes = self._swa_slot_bytes() - compressed_cell = self._compressed_cell_size() - - if self.max_request_num is not None and self.sliding_window is not None and compressed_cell > 0: - swa_budget = int(self.max_request_num) * self._swa_per_req_budget() + self.swa_extra_token_num - full_cell = swa_slot_bytes + compressed_cell - if available_bytes <= full_cell * swa_budget: - # 小显存: full token 数还到不了 swa 预算,swa 池跟随 full token 数(每 token 一个 swa 槽)。 - self.size = max(1, int(available_bytes / full_cell)) - else: - size_budget = max(1, int((available_bytes - swa_slot_bytes * swa_budget) / compressed_cell)) - if size_budget * self.swa_full_tokens_ratio > swa_budget: - # 比例下限生效(_planned_swa_size 会取 ratio*full),按该机制反解 full。 - self.size = max( - 1, int(available_bytes / (swa_slot_bytes * self.swa_full_tokens_ratio + compressed_cell)) - ) - else: - self.size = size_budget - else: - self.size = max(1, int(available_bytes / (swa_slot_bytes + compressed_cell))) - - if world_size > 1: - tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") - dist.all_reduce(tensor, op=dist.ReduceOp.MIN) - self.size = tensor.item() - - logger.info( - f"{str(available_memory)} GB space is available after load the model weight\n" - f"{str(self.get_cell_size() / 1024 ** 2)} MB is the conservative size of one token kv cache\n" - f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" + kv_bytes = self.mla_bytes_per_token + indexer_bytes = self.indexer_bytes_per_token + state_dtype_bytes = torch._utils._element_size(torch.float32) + c4_state_width = 4 * self.head_dim + 4 * self.indexer_head_dim + c128_state_width = 2 * self.head_dim + c4_state_bytes = DSV4_C4_STATE_RING / DSV4_SWA_PAGE_SIZE * c4_state_width * state_dtype_bytes * self.n_c4 + c128_state_bytes = ( + DSV4_C128_STATE_RING / DSV4_SWA_PAGE_SIZE * c128_state_width * state_dtype_bytes * self.n_c128 ) - return + swa_slot = kv_bytes * self.layer_num + c4_state_bytes + c128_state_bytes + compressed = (kv_bytes + indexer_bytes) * self.n_c4 / 4 + kv_bytes * self.n_c128 / 128 + + return swa_slot * self.swa_full_tokens_ratio + compressed # ------------------------------------------------------------------ buffers def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): @@ -305,7 +225,6 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): server = get_unique_server_name() self.swa_size = self._planned_swa_size(size) - assert self.swa_size % DSV4_SWA_PAGE_SIZE == 0 self.swa_pool = PackedPagePool( size=self.swa_size, page_size=DSV4_SWA_PAGE_SIZE, @@ -777,7 +696,7 @@ def pack_mla_kv_to_cache_fused_norm_rope( 省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" if kv.shape[0] == 0: return - from sglang.jit_kernel.dsv4 import fused_k_norm_rope_flashmla + from lightllm.third_party.sglang_jit.dsv4 import fused_k_norm_rope_flashmla swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] # 未映射槽位(-1, 如 decode 图 warmup 的 HOLD 行: prep 跳过 alloc_swa)对老 triton diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 5e7c4f96dd..18ef51dbe3 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -398,16 +398,6 @@ def __init__( ): super().__init__(max_request_num, max_sequence_length, mem_manager) - self.mem_manager = mem_manager - if mem_manager is not None: - if compress_rates is None: - compress_rates = mem_manager.compress_rates - if head_dim is None: - head_dim = mem_manager.head_dim - if indexer_head_dim is None: - indexer_head_dim = mem_manager.indexer_head_dim - if sliding_window is None: - sliding_window = mem_manager.sliding_window self.sliding_window = sliding_window # 出窗回收水位线: -1 表示该 req 尚未见过 prefill chunk(首个 chunk 的 ready_cache_len # 即共享前缀边界,作为永不下探的回收下界)。 @@ -430,18 +420,6 @@ def __init__( return - def bind_mem_manager(self, mem_manager: DeepseekV4MemoryManager): - assert isinstance(mem_manager, DeepseekV4MemoryManager) - assert self.compress_rates == mem_manager.compress_rates - assert self.head_dim == mem_manager.head_dim - assert self.indexer_head_dim == mem_manager.indexer_head_dim - if self.sliding_window is None: - self.sliding_window = mem_manager.sliding_window - else: - assert mem_manager.sliding_window is None or self.sliding_window == mem_manager.sliding_window - self.mem_manager = mem_manager - return - # ------------------------------------------------------------------ swa slot prep (per step) def _swa_retain_len(self) -> int: """出窗回收的保留长度 = window + 一个 radix 页。 diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 887e72f433..c824b24387 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -85,9 +85,6 @@ def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() compress_rates = getattr(self, "_dsv4_compress_rates", self._get_compress_rates(layer_num)) sliding_window = int(self.config["sliding_window"]) - # 活跃窗口之外的 swa 余量: 在途 prefill chunk 的瞬时占用(出窗槽位到下一次 prep 才回收) - # + radix cache 持有的窗口尾部(每条缓存序列约一个 window)。 - swa_extra_token_num = int(self.batch_max_tokens or 0) + self.max_req_num * sliding_window self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, @@ -98,11 +95,9 @@ def _init_mem_manager(self): indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, sliding_window=sliding_window, - swa_extra_token_num=swa_extra_token_num, mem_fraction=self.mem_fraction, ) - assert isinstance(self.req_manager, DeepseekV4ReqManager) - self.req_manager.bind_mem_manager(self.mem_manager) + self.req_manager.mem_manager = self.mem_manager return def _init_cudagraph(self): From cf433fb73424333c07cba180c1c650dd3efd9401 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 01:24:36 +0000 Subject: [PATCH 032/214] fix swa insufficient --- .../router/dynamic_prompt/radix_cache.py | 36 +++++++++++++++ .../server/router/model_infer/infer_batch.py | 46 +++++++++++++++++++ .../model_infer/mode_backend/base_backend.py | 17 ++++++- 3 files changed, 97 insertions(+), 2 deletions(-) diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index ff09c018be..ececce7b27 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -147,6 +147,28 @@ def __init__( f"{unique_name}_tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64 ) self.tree_total_tokens_num.arr[0] = 0 + self.swa_tree_total_pages_num = 0 + self.swa_refed_pages_num = 0 + # 每个 prompt-cache 页折算多少 swa 页(DSV4 为 256/128=2);非 swa 场景为 0,_node_swa_pages_num 退化为常数 0。 + self._swa_pages_per_prompt_page = self._probe_swa_pages_per_prompt_page() + + def _probe_swa_pages_per_prompt_page(self) -> int: + """构造期探测一次 mem_manager 是否带 swa_pool,缓存折算系数,避免热路径反复 getattr。""" + if self.mem_manager is None or self.extra_value_ops is None: + return 0 + swa_pool = getattr(self.mem_manager, "swa_pool", None) + swa_page_size = getattr(swa_pool, "page_size", None) + if swa_page_size is None: + return 0 + return (self.page_size + int(swa_page_size) - 1) // int(swa_page_size) + + def _node_swa_pages_num(self, node: TreeNode) -> int: + if self._swa_pages_per_prompt_page == 0 or node.token_extra_value is None: + return 0 + valid = node.token_extra_value.swa_page_valid + if valid is None: + return 0 + return int(valid.sum().item()) * self._swa_pages_per_prompt_page def _align_len(self, length: int) -> int: if self.page_size <= 1: @@ -281,6 +303,7 @@ def _insert_helper_no_recursion( ) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) + self.swa_tree_total_pages_num += self._node_swa_pages_num(new_node) if split_parent_node.is_leaf(): self.evict_tree_set.add(split_parent_node) @@ -310,6 +333,7 @@ def _insert_helper_no_recursion( ) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) + self.swa_tree_total_pages_num += self._node_swa_pages_num(new_node) if new_node.is_leaf(): self.evict_tree_set.add(new_node) return 0, new_node @@ -399,6 +423,7 @@ def _match_prefix_helper_no_recursion( # from 0 to 1 need update refs token num if node.ref_counter == 1: self.refed_tokens_num.arr[0] += len(node.token_mem_index_value) + self.swa_refed_pages_num += self._node_swa_pages_num(node) if len(key) == 0: return node @@ -428,6 +453,7 @@ def _match_prefix_helper_no_recursion( # from 0 to 1 need update refs token num if split_parent_node.ref_counter == 1: self.refed_tokens_num.arr[0] += len(split_parent_node.token_mem_index_value) + self.swa_refed_pages_num += self._node_swa_pages_num(split_parent_node) if child.is_leaf(): self.evict_tree_set.add(child) @@ -455,6 +481,7 @@ def evict(self, need_remove_tokens, evict_callback): self.extra_value_ops.free(node.token_extra_value) # update total token num self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value) + self.swa_tree_total_pages_num -= self._node_swa_pages_num(node) parent_node: TreeNode = node.parent parent_node.remove_child(node) if parent_node.is_leaf(): @@ -537,6 +564,8 @@ def clear_tree_nodes(self): self.tree_total_tokens_num.arr[0] = 0 self.refed_tokens_num.arr[0] = 0 + self.swa_tree_total_pages_num = 0 + self.swa_refed_pages_num = 0 return def dec_node_ref_counter(self, node: TreeNode): @@ -550,6 +579,7 @@ def dec_node_ref_counter(self, node: TreeNode): while node is not None: if node.ref_counter == 1: self.refed_tokens_num.arr[0] -= len(node.token_mem_index_value) + self.swa_refed_pages_num -= self._node_swa_pages_num(node) node.ref_counter -= 1 node = node.parent @@ -569,6 +599,7 @@ def add_node_ref_counter(self, node: TreeNode): while node is not None: if node.ref_counter == 0: self.refed_tokens_num.arr[0] += len(node.token_mem_index_value) + self.swa_refed_pages_num += self._node_swa_pages_num(node) node.ref_counter += 1 node = node.parent @@ -608,6 +639,9 @@ def get_refed_tokens_num(self): def get_tree_total_tokens_num(self): return self.tree_total_tokens_num.arr[0] + def get_unrefed_swa_pages_num(self): + return self.swa_tree_total_pages_num - self.swa_refed_pages_num + def print_self(self, indent=0): self._print_helper(self.root_node, indent) @@ -644,9 +678,11 @@ def reclaim_unreferenced_swa_pages(self, need_pages: int) -> None: # 每回收一个节点就复查目标,避免多回收(无谓削减命中可用性)。 while node is not None and node is not self.root_node and node.ref_counter == 0: if len(node.token_mem_index_value) > 0: + old_pages = self._node_swa_pages_num(node) self.mem_manager.evict_swa(node.token_mem_index_value) if node.token_extra_value is not None: invalidate(node.token_extra_value) + self.swa_tree_total_pages_num -= old_pages - self._node_swa_pages_num(node) if allocator.can_use_mem_size >= target: return node = node.parent diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 2bf2314185..b04c09a056 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1,4 +1,5 @@ import enum +from functools import lru_cache import torch import torch.distributed as dist import numpy as np @@ -437,6 +438,51 @@ def get_can_alloc_token_num(self): ) return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num + def get_can_alloc_dsv4_swa_page_num(self): + mem_manager = self.req_manager.mem_manager + allocator = getattr(mem_manager, "swa_page_allocator", None) + if allocator is None: + return None + + radix_cache_unref_page_num = 0 + if self.radix_cache is not None: + radix_cache_unref_page_num = self.radix_cache.get_unrefed_swa_pages_num() + return int(allocator.can_use_mem_size) + radix_cache_unref_page_num + + def get_dsv4_swa_prefill_need_page_num(self, req: "InferReq", is_chuncked_prefill: bool): + page_size = self._get_dsv4_swa_page_size() + if page_size is None: + return 0 + + start = int(req.cur_kv_len) + if is_chuncked_prefill: + end = int(req.get_chuncked_input_token_len()) + else: + end = int(req.get_cur_total_len()) + if end <= start: + return 0 + first_new_page = (start + page_size - 1) // page_size + last_page = (end - 1) // page_size + return last_page - first_new_page + 1 + + def get_dsv4_swa_decode_need_page_num(self, req: "InferReq"): + page_size = self._get_dsv4_swa_page_size() + if page_size is None: + return 0 + + seq_len = int(req.get_cur_total_len()) + if seq_len <= 0: + return 0 + return 1 if (seq_len - 1) % page_size == 0 else 0 + + @lru_cache + def _get_dsv4_swa_page_size(self): + mem_manager = self.req_manager.mem_manager + allocator = getattr(mem_manager, "swa_page_allocator", None) + if allocator is None: + return None + return mem_manager.swa_pool.page_size + def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """ 该函数用于在线性混合模型prefill后,如果存在大页匹配的情况下,将线性层状态复制到 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 1f19b9462a..ceb542e541 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -597,6 +597,7 @@ def _get_classed_reqs( prefill_tokens = 0 can_alloc_token_num = g_infer_context.get_can_alloc_token_num() + can_alloc_dsv4_swa_page_num = g_infer_context.get_can_alloc_dsv4_swa_page_num() for req_obj in ready_reqs: @@ -630,9 +631,14 @@ def _get_classed_reqs( if is_decode: token_num = req_obj.decode_need_token_num() - if token_num <= can_alloc_token_num: + swa_page_num = g_infer_context.get_dsv4_swa_decode_need_page_num(req_obj) + if token_num <= can_alloc_token_num and ( + can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num + ): decode_reqs.append(req_obj) can_alloc_token_num -= token_num + if can_alloc_dsv4_swa_page_num is not None: + can_alloc_dsv4_swa_page_num -= swa_page_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True @@ -647,10 +653,17 @@ def _get_classed_reqs( token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill) if prefill_tokens + token_num > self.batch_max_tokens: continue - if token_num <= can_alloc_token_num: + swa_page_num = g_infer_context.get_dsv4_swa_prefill_need_page_num( + req_obj, is_chuncked_prefill=not self.disable_chunked_prefill + ) + if token_num <= can_alloc_token_num and ( + can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num + ): prefill_tokens += token_num prefill_reqs.append(req_obj) can_alloc_token_num -= token_num + if can_alloc_dsv4_swa_page_num is not None: + can_alloc_dsv4_swa_page_num -= swa_page_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True From 40f5810125972b2555e41d9435ab8bbe7abbc056 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 01:56:23 +0000 Subject: [PATCH 033/214] fix --- lightllm/server/router/model_infer/infer_batch.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index b04c09a056..f73fe6cad4 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1,5 +1,4 @@ import enum -from functools import lru_cache import torch import torch.distributed as dist import numpy as np @@ -475,7 +474,6 @@ def get_dsv4_swa_decode_need_page_num(self, req: "InferReq"): return 0 return 1 if (seq_len - 1) % page_size == 0 else 0 - @lru_cache def _get_dsv4_swa_page_size(self): mem_manager = self.req_manager.mem_manager allocator = getattr(mem_manager, "swa_page_allocator", None) From f527ca25db4f1cd3795863b628dd6f01224313a1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 03:07:03 +0000 Subject: [PATCH 034/214] rename --- .../deepseek4_mem_manager.py | 23 ++++++++----------- .../router/dynamic_prompt/radix_cache.py | 6 ++--- .../model_infer/mode_backend/base_backend.py | 6 ++--- 3 files changed, 16 insertions(+), 19 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 8dedf16387..314b556a8a 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -120,7 +120,7 @@ class DeepseekV4MemoryManager(MemoryManager): 映射记录到 ``full_to_swa_indexs``(以 full token 槽位为键)。出窗槽位由 DeepseekV4ReqManager 在 prep 阶段批量惰性回收(``evict_swa``,页存活计数减到 0 才整页归还);full 槽位释放时 ``free`` 级联回收对应 swa 槽,所以 radix 驱逐/请求释放/暂停无需任何额外协议。 - 页 allocator 触底时先走压力阀(radix 对 ref==0 节点回收)再 assert。 + 页 allocator 触底时先走 swa free hook(radix 对 ref==0 节点 free)再 assert。 没有 ring buffer,prefill chunk 大小不受 sliding_window 限制。 - ``c4_pool``/``c128_pool``: 压缩 latent,按 qwen3next 的层号压实手法只为压缩层建层; c4 另带 packed indexer-K 池。槽位映射(``full_to_c4/c128_indexs``)以组末 token 的 full @@ -187,10 +187,7 @@ def __init__( # ------------------------------------------------------------------ sizing def _planned_swa_size(self, full_size: int) -> int: - # swa 池 = full * ratio,按 SWA 物理页(128)向上取整;与 sglang DSV4PoolConfigurator 同思路。 - cap = int(full_size * self.swa_full_tokens_ratio) - cap = max(1, min(full_size, cap)) - return _ceil_div(cap, DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE + return _ceil_div(int(full_size * self.swa_full_tokens_ratio), DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE @staticmethod def _paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: @@ -245,9 +242,9 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): # 页存活计数 = 指向该页的有效 full_to_swa 行数;减到 0 归还 allocator(出窗逐 token # 回收下,「部分出窗页」计数 > 0 自然受保护)。下标含 HOLD 页(只读不增减)。 self.swa_page_live_count = torch.zeros((self.swa_pool.num_pages,), dtype=torch.int32, device="cuda") - # swa 压力阀(可选): 页 allocator 触底时回调(radix 对 ref==0 节点回收 swa 页), - # 由 backend 在 radix cache 创建后注入;assert 仍是最后防线。 - self._swa_pressure_valve = None + # swa free hook(可选): 页 allocator 触底时回调(radix 对 ref==0 节点 free swa 页), + # 由 backend 在 radix cache 创建后 register;assert 仍是最后防线。 + self._free_radix_unreferenced_swa_fn = None self.full_to_swa_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") self.full_to_swa_indexs[size] = self.swa_pool.HOLD_TOKEN_MEMINDEX @@ -365,14 +362,14 @@ def get_c128_state_buffer(self, layer_index: int) -> torch.Tensor: return self.c128_state_buffer[self.layer_to_c128_idx[layer_index]] # ------------------------------------------------------------------ swa slot lifecycle - def set_swa_pressure_valve(self, valve) -> None: - """valve(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点回收 swa)。""" - self._swa_pressure_valve = valve + def register_swa_free_hook(self, fn) -> None: + """fn(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点 free swa)。""" + self._free_radix_unreferenced_swa_fn = fn return def _alloc_swa_pages(self, need_pages: int) -> torch.Tensor: - if need_pages > self.swa_page_allocator.can_use_mem_size and self._swa_pressure_valve is not None: - self._swa_pressure_valve(need_pages - self.swa_page_allocator.can_use_mem_size) + if need_pages > self.swa_page_allocator.can_use_mem_size and self._free_radix_unreferenced_swa_fn is not None: + self._free_radix_unreferenced_swa_fn(need_pages - self.swa_page_allocator.can_use_mem_size) return self.swa_page_allocator.alloc(need_pages) def _count_swa_pages(self, swa_slots: torch.Tensor, delta: int) -> torch.Tensor: diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index ececce7b27..dbffea0a33 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -656,9 +656,9 @@ def _print_helper(self, node: TreeNode, indent): self._print_helper(child, indent=indent + 2) return - def reclaim_unreferenced_swa_pages(self, need_pages: int) -> None: - """DeepSeek-V4 swa 压力阀: 页 allocator 触底时,沿 LRU 序(evict_tree_set)只对 - ref_count==0 的节点链回收其 swa 页(full 槽与压缩条目保留——节点仍可服务更长前缀的 + def free_unreferenced_swa_pages(self, need_pages: int) -> None: + """DeepSeek-V4 swa free hook: 页 allocator 触底时,沿 LRU 序(evict_tree_set)只对 + ref_count==0 的节点链 free 其 swa 页(full 槽与压缩条目保留——节点仍可服务更长前缀的 中段命中),并清载荷 bitmap 位使后续命中按缩短语义裁剪。所有权判定直接复用 radix 引用计数: 节点被任何活跃请求借用即 ref>0,其页不可达。不够时由 allocator 的 assert 兜底(最后防线)。""" 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 ceb542e541..3319a42db5 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -194,9 +194,9 @@ def init_model(self, kvargs): page_size=radix_page_size, extra_value_ops=radix_extra_value_ops, ) - if radix_extra_value_ops is not None and hasattr(self.model.mem_manager, "set_swa_pressure_valve"): - # swa 页 allocator 触底时让 radix 对 ref==0 节点回收 swa 页(DeepSeek-V4)。 - self.model.mem_manager.set_swa_pressure_valve(self.radix_cache.reclaim_unreferenced_swa_pages) + if radix_extra_value_ops is not None and hasattr(self.model.mem_manager, "register_swa_free_hook"): + # swa 页 allocator 触底时让 radix 对 ref==0 节点 free swa 页(DeepSeek-V4)。 + self.model.mem_manager.register_swa_free_hook(self.radix_cache.free_unreferenced_swa_pages) if "prompt_cache_kv_buffer" in model_cfg: assert self.use_dynamic_prompt_cache From 255e90d01a3e08226d7df9a181a153a32336fff7 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 09:24:46 +0000 Subject: [PATCH 035/214] tune config --- ..._num=1,use_fp8_w8a8=true}_NVIDIA_H200.json | 110 +++++++++++++ ..._num=6,use_fp8_w8a8=true}_NVIDIA_H200.json | 110 +++++++++++++ .../{topk_num=6}_NVIDIA_H200.json | 50 ++++++ ...orch.bfloat16,topk_num=6}_NVIDIA_H200.json | 74 +++++++++ ...M=8,dtype=torch.bfloat16}_NVIDIA_H200.json | 74 +++++++++ ...out_dtype=torch.bfloat16}_NVIDIA_H200.json | 146 ++++++++++++++++++ 6 files changed, 564 insertions(+) create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json new file mode 100644 index 0000000000..b1aae6bfba --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json @@ -0,0 +1,110 @@ +{ + "12288": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "1536": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "192": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "24576": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "384": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "48": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "49152": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "6": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "600": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "6144": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "768": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "96": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json new file mode 100644 index 0000000000..9ffb0efd19 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json @@ -0,0 +1,110 @@ +{ + "1": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "100": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "1024": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "128": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 5, + "num_warps": 4 + }, + "16": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "2048": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "256": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "32": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 4, + "num_warps": 4 + }, + "4096": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": false, + "num_stages": 5, + "num_warps": 8 + }, + "64": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 5, + "num_warps": 4 + }, + "8": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 4, + "num_warps": 4 + }, + "8192": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json new file mode 100644 index 0000000000..85a20d9b1b --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json @@ -0,0 +1,50 @@ +{ + "1": { + "BLOCK_SIZE": 128, + "num_warps": 1 + }, + "100": { + "BLOCK_SIZE": 128, + "num_warps": 4 + }, + "1024": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "128": { + "BLOCK_SIZE": 128, + "num_warps": 4 + }, + "16": { + "BLOCK_SIZE": 256, + "num_warps": 8 + }, + "2048": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "256": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "32": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "4096": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "64": { + "BLOCK_SIZE": 128, + "num_warps": 4 + }, + "8": { + "BLOCK_SIZE": 512, + "num_warps": 4 + }, + "8192": { + "BLOCK_SIZE": 256, + "num_warps": 8 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json new file mode 100644 index 0000000000..de2f015a04 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_DIM": 128, + "BLOCK_M": 1, + "NUM_STAGE": 2, + "num_warps": 4 + }, + "100": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "1024": { + "BLOCK_DIM": 256, + "BLOCK_M": 4, + "NUM_STAGE": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + }, + "16": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + }, + "2048": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 4, + "num_warps": 2 + }, + "256": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 2, + "num_warps": 1 + }, + "32": { + "BLOCK_DIM": 256, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 2 + }, + "4096": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + }, + "8": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "8192": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 4, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json new file mode 100644 index 0000000000..88742a0b13 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 4, + "num_warps": 1 + }, + "100": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 2 + }, + "1024": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 1, + "num_warps": 4 + }, + "128": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 2 + }, + "16": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 4 + }, + "2048": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 1, + "num_warps": 2 + }, + "256": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 2 + }, + "32": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 8 + }, + "4096": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 2, + "num_warps": 1 + }, + "64": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 1, + "num_warps": 4 + }, + "8": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 2, + "num_warps": 8 + }, + "8192": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 2, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json new file mode 100644 index 0000000000..fbd3649737 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json @@ -0,0 +1,146 @@ +{ + "1": { + "BLOCK_M": 128, + "BLOCK_N": 32, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "100": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "1024": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "12288": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "128": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1536": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "16": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "192": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "2048": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "24576": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "256": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "32": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "384": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "4096": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "48": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "49152": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "6": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "600": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "6144": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "64": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "768": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "8": { + "BLOCK_M": 1, + "BLOCK_N": 64, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "8192": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "96": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + } +} \ No newline at end of file From d88dc71c99b1f053c591a410ea2fcbf546c2796a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 14:10:42 +0000 Subject: [PATCH 036/214] prepare opt --- lightllm/common/basemodel/basemodel.py | 73 +++++++++++------- lightllm/common/basemodel/batch_objs.py | 16 ++++ .../deepseek4_mem_manager.py | 10 ++- lightllm/common/req_manager.py | 74 ++++++++++++++++--- .../mode_backend/chunked_prefill/impl.py | 1 + .../mode_backend/dp_backend/impl.py | 3 + 6 files changed, 137 insertions(+), 40 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 5825d9b45f..f440b98213 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -291,11 +291,15 @@ def _init_custom(self): @torch.no_grad() def forward(self, model_input: ModelInput): - # decode 槽位 prep: 放在 to_cuda 前 (b_req_idx/b_seq_len/mem_indexes_cpu 还是原生 CPU 张量), + # decode 槽位 prep: 放在 to_cuda 前, 优先使用 b_req_idx/b_seq_len 的 CPU mirror, # 且此刻已在 forward 的 CUDA stream 上 -> 与后续 attention 同流, 无跨流竞态、无 D2H。 # mem_indexes_cpu is None 时跳过: cudagraph warmup 的输入全在 CUDA 且 b_req_idx 全为 HOLD, prep 本就是 no-op。 if not model_input.is_prefill and model_input.mem_indexes_cpu is not None: - self.req_manager.prepare_decode(model_input.b_req_idx, model_input.b_seq_len, model_input.mem_indexes_cpu) + self.req_manager.prepare_decode( + model_input.b_req_idx_cpu, + model_input.b_seq_len_cpu, + model_input.mem_indexes_cpu, + ) model_input.to_cuda() assert model_input.mem_indexes.is_cuda @@ -376,6 +380,15 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mtp_index, (0, padded_batch_size), mode="constant", value=0 ) new_model_input.b_seq_len = F.pad(new_model_input.b_seq_len, (0, padded_batch_size), mode="constant", value=2) + new_model_input.b_req_idx_cpu = F.pad( + new_model_input.b_req_idx_cpu, + (0, padded_batch_size), + mode="constant", + value=self.req_manager.HOLD_REQUEST_ID, + ) + new_model_input.b_seq_len_cpu = F.pad( + new_model_input.b_seq_len_cpu, (0, padded_batch_size), mode="constant", value=2 + ) new_model_input.mem_indexes = F.pad( new_model_input.mem_indexes, (0, padded_batch_size), @@ -433,6 +446,15 @@ 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_req_idx_cpu = F.pad( + new_model_input.b_req_idx_cpu, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID + ) + new_model_input.b_seq_len_cpu = F.pad( + new_model_input.b_seq_len_cpu, (0, 1), mode="constant", value=padded_token_num + ) + new_model_input.b_ready_cache_len_cpu = F.pad( + new_model_input.b_ready_cache_len_cpu, (0, 1), mode="constant", value=0 + ) 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 可能会让外面使用的数组引用发生变化,导致错误。 @@ -526,17 +548,14 @@ def _prefill( alloc_mem_index=infer_state.mem_index, max_q_seq_len=infer_state.max_q_seq_len, ) - if hasattr(self.req_manager, "prepare_prefill_swa"): - self.req_manager.prepare_prefill_swa( - b_req_idx=infer_state.b_req_idx, - b_ready_cache_len=infer_state.b_ready_cache_len, - b_seq_len=infer_state.b_seq_len, - ) - if hasattr(self.req_manager, "prepare_prefill_compress_slots"): - self.req_manager.prepare_prefill_compress_slots( + if model_input.b_req_idx_cpu is not None: + self.req_manager.prepare_prefill( b_req_idx=infer_state.b_req_idx, b_ready_cache_len=infer_state.b_ready_cache_len, b_seq_len=infer_state.b_seq_len, + b_req_idx_cpu=model_input.b_req_idx_cpu, + b_ready_cache_len_cpu=model_input.b_ready_cache_len_cpu, + b_seq_len_cpu=model_input.b_seq_len_cpu, ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -758,17 +777,14 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state0.mem_index, max_q_seq_len=infer_state0.max_q_seq_len, ) - if hasattr(self.req_manager, "prepare_prefill_swa"): - self.req_manager.prepare_prefill_swa( - b_req_idx=infer_state0.b_req_idx, - b_ready_cache_len=infer_state0.b_ready_cache_len, - b_seq_len=infer_state0.b_seq_len, - ) - if hasattr(self.req_manager, "prepare_prefill_compress_slots"): - self.req_manager.prepare_prefill_compress_slots( + if model_input0.b_req_idx_cpu is not None: + self.req_manager.prepare_prefill( b_req_idx=infer_state0.b_req_idx, b_ready_cache_len=infer_state0.b_ready_cache_len, b_seq_len=infer_state0.b_seq_len, + b_req_idx_cpu=model_input0.b_req_idx_cpu, + b_ready_cache_len_cpu=model_input0.b_ready_cache_len_cpu, + b_seq_len_cpu=model_input0.b_seq_len_cpu, ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -783,17 +799,14 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state1.mem_index, max_q_seq_len=infer_state1.max_q_seq_len, ) - if hasattr(self.req_manager, "prepare_prefill_swa"): - self.req_manager.prepare_prefill_swa( - b_req_idx=infer_state1.b_req_idx, - b_ready_cache_len=infer_state1.b_ready_cache_len, - b_seq_len=infer_state1.b_seq_len, - ) - if hasattr(self.req_manager, "prepare_prefill_compress_slots"): - self.req_manager.prepare_prefill_compress_slots( + if model_input1.b_req_idx_cpu is not None: + self.req_manager.prepare_prefill( b_req_idx=infer_state1.b_req_idx, b_ready_cache_len=infer_state1.b_ready_cache_len, b_seq_len=infer_state1.b_seq_len, + b_req_idx_cpu=model_input1.b_req_idx_cpu, + b_ready_cache_len_cpu=model_input1.b_ready_cache_len_cpu, + b_seq_len_cpu=model_input1.b_seq_len_cpu, ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -822,10 +835,14 @@ 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): - # decode 槽位 prep: 在 to_cuda 前 (原生 CPU 张量)、且已在 forward 的 CUDA stream 上 (见 forward 注释)。 + # decode 槽位 prep: 在 to_cuda 前使用 CPU mirror, 且已在 forward 的 CUDA stream 上 (见 forward 注释)。 for mi in (model_input0, model_input1): if mi.mem_indexes_cpu is not None: - self.req_manager.prepare_decode(mi.b_req_idx, mi.b_seq_len, mi.mem_indexes_cpu) + self.req_manager.prepare_decode( + mi.b_req_idx_cpu, + mi.b_seq_len_cpu, + mi.mem_indexes_cpu, + ) model_input0.to_cuda() model_input1.to_cuda() assert self.args.enable_tpsp_mix_mode diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 1795ff9a82..7110d4895a 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -42,6 +42,9 @@ class ModelInput: multimodal_params: list = None # cpu 变量 mem_indexes_cpu: torch.Tensor = None + b_req_idx_cpu: torch.Tensor = None + b_seq_len_cpu: torch.Tensor = None + b_ready_cache_len_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 # 的一些变量 b_prefill_has_output_cpu: List[bool] = None # 标记进行prefill的请求是否具有输出 @@ -53,6 +56,18 @@ class ModelInput: # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None + def _capture_cpu_mirror(self, tensor_name: str, mirror_name: str): + tensor = getattr(self, tensor_name) + if tensor is not None and not tensor.is_cuda: + setattr(self, mirror_name, tensor) + return + + def capture_cpu_mirrors(self): + self._capture_cpu_mirror("b_req_idx", "b_req_idx_cpu") + self._capture_cpu_mirror("b_seq_len", "b_seq_len_cpu") + self._capture_cpu_mirror("b_ready_cache_len", "b_ready_cache_len_cpu") + return + def to_cuda(self): if self.input_ids is not None: self.input_ids = self.input_ids.cuda(non_blocking=True) @@ -82,6 +97,7 @@ def to_cuda(self): self.b_shared_seq_len = self.b_shared_seq_len.cuda(non_blocking=True) def __post_init__(self): + self.capture_cpu_mirrors() self.check_input() def check_input(self): diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 314b556a8a..784065d964 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -385,6 +385,9 @@ def alloc_swa_prefill( b_ready_cache_len: torch.Tensor, b_seq_len: torch.Tensor, req_to_token_indexs: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_ready_cache_len_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, ) -> None: """prefill prep: 为各请求位置 [ready, seq) 的新 token 分配位置对齐的 swa 槽。 @@ -396,9 +399,10 @@ def alloc_swa_prefill( """ page = DSV4_SWA_PAGE_SIZE hold_req_id = self.max_request_num # padding 行的请求 id(req_manager.HOLD_REQUEST_ID) - req_list = b_req_idx.detach().cpu().tolist() - ready_list = b_ready_cache_len.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() + + req_list = b_req_idx_cpu.tolist() + ready_list = b_ready_cache_len_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() segs = [] # (req_idx, start, end, n_new_pages, has_cont_page) total_new_pages = 0 diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 18ef51dbe3..e633f9ee3d 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -161,8 +161,21 @@ def free_all(self): self.req_list = _ReqLinkedList(self.max_request_num) return + def prepare_prefill( + self, + b_req_idx, + b_ready_cache_len, + b_seq_len, + b_req_idx_cpu=None, + b_ready_cache_len_cpu=None, + b_seq_len_cpu=None, + ): + """prefill 在 init_req_to_token_indexes 之后调用的钩子。基类 no-op; 需要 + prefill KV 槽位 prep 的模型 (DeepSeek-V4) override。""" + return + def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): - """每个 decode step 在 to_cuda 之前调用的钩子 (数据为原生 CPU 张量, 且已在 forward 的 + """每个 decode step 在 to_cuda 之前调用的钩子 (优先使用 CPU mirror, 且已在 forward 的 CUDA stream 上)。基类 no-op; 需要 per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" return @@ -375,7 +388,7 @@ class DeepseekV4ReqManager(ReqManager): 真实 mem_manager 会在 `_init_mem_manager()` 后通过 `bind_mem_manager()` 接入。 * 压缩槽位不在本类: ``full_to_c4/c128_indexs``(mem manager)以组末 token 的 full 槽位为键。 - 本类只负责 prep 阶段的分配与 scatter(``prepare_prefill_compress_slots`` / + 本类只负责 prep 阶段的分配与 scatter(``prepare_prefill`` / ``prepare_decode_compress_slots``)——必须先于 attention metadata 构建/图捕获; 条目内容由 layer-infer 的 compressor 前向写入。 * compressor 在途状态不在本类: c4/c128 都在 mem manager 的 swa 页派生池, @@ -434,6 +447,9 @@ def prepare_prefill_swa( b_req_idx: torch.Tensor, b_ready_cache_len: torch.Tensor, b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_ready_cache_len_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, ) -> None: """prefill prep: 为本 chunk 全部新 token(位置 [ready, seq))分配位置对齐的 swa 槽, 并回收已出窗位置的槽。 @@ -445,8 +461,8 @@ def prepare_prefill_swa( if self.sliding_window is not None: retain = self._swa_retain_len() evict_slots = [] - req_list = b_req_idx.detach().cpu().tolist() - ready_list = b_ready_cache_len.detach().cpu().tolist() + req_list = b_req_idx_cpu.tolist() + ready_list = b_ready_cache_len_cpu.tolist() for req_idx, ready_len in zip(req_list, ready_list): req_idx = int(req_idx) if req_idx == self.HOLD_REQUEST_ID: @@ -463,16 +479,53 @@ def prepare_prefill_swa( self._swa_evict_marks[req_idx] = evict_end if evict_slots: self.mem_manager.evict_swa(torch.cat(evict_slots)) - self.mem_manager.alloc_swa_prefill(b_req_idx, b_ready_cache_len, b_seq_len, self.req_to_token_indexs) + self.mem_manager.alloc_swa_prefill( + b_req_idx, + b_ready_cache_len, + b_seq_len, + self.req_to_token_indexs, + b_req_idx_cpu=b_req_idx_cpu, + b_ready_cache_len_cpu=b_ready_cache_len_cpu, + b_seq_len_cpu=b_seq_len_cpu, + ) return def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): """decode 每步槽位 prep: 先 swa 再 compress。由 BaseModel.forward / microbatch_overlap_decode - 在 to_cuda 之前调用 (CPU 数据 + forward 的 CUDA stream); 不再放在 _decode 里。""" + 在 to_cuda 之前调用 (CPU mirror + forward 的 CUDA stream); 不再放在 _decode 里。""" self.prepare_decode_swa(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) self.prepare_decode_compress_slots(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) return + def prepare_prefill( + self, + b_req_idx: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_ready_cache_len_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, + ) -> None: + """prefill 槽位 prep: 先 swa 再 compress。由 BaseModel 在 + init_req_to_token_indexes 之后、attention metadata 构建之前调用。""" + self.prepare_prefill_swa( + b_req_idx, + b_ready_cache_len, + b_seq_len, + b_req_idx_cpu=b_req_idx_cpu, + b_ready_cache_len_cpu=b_ready_cache_len_cpu, + b_seq_len_cpu=b_seq_len_cpu, + ) + self.prepare_prefill_compress_slots( + b_req_idx, + b_ready_cache_len, + b_seq_len, + b_req_idx_cpu=b_req_idx_cpu, + b_ready_cache_len_cpu=b_ready_cache_len_cpu, + b_seq_len_cpu=b_seq_len_cpu, + ) + return + def prepare_decode_swa( self, b_req_idx_cpu: torch.Tensor, @@ -663,15 +716,18 @@ def prepare_prefill_compress_slots( b_req_idx: torch.Tensor, b_ready_cache_len: torch.Tensor, b_seq_len: torch.Tensor, + b_req_idx_cpu: torch.Tensor, + b_ready_cache_len_cpu: torch.Tensor, + b_seq_len_cpu: torch.Tensor, ) -> None: """prefill prep: 为本 chunk 内的组末 token(位置 (g+1)*ratio-1 ∈ [ready, seq))分配压缩槽, scatter 进 full_to_c4/c128_indexs。必须在 init_req_to_token_indexes 之后(组末 full 槽 从 req_to_token_indexs 取)、attention metadata 构建之前调用。""" if self.n_c4 == 0 and self.n_c128 == 0: return - req_list = b_req_idx.detach().cpu().tolist() - ready_list = b_ready_cache_len.detach().cpu().tolist() - seq_list = b_seq_len.detach().cpu().tolist() + req_list = b_req_idx_cpu.tolist() + ready_list = b_ready_cache_len_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() if self.n_c4 > 0: for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): req_idx = int(req_idx) 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 792a10a788..40d1c175aa 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 @@ -402,6 +402,7 @@ def _draft_decode_eagle( 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.b_seq_len_cpu += 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( 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 e6b9d1c18d..b051a1a3ab 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 @@ -603,6 +603,7 @@ def _draft_decode_eagle( 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.b_seq_len_cpu += 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( @@ -967,6 +968,7 @@ def _draft_decode_eagle_overlap( ) draft_model_input0.b_seq_len += 1 + draft_model_input0.b_seq_len_cpu += 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( @@ -981,6 +983,7 @@ def _draft_decode_eagle_overlap( ).view(-1) draft_model_input1.b_seq_len += 1 + draft_model_input1.b_seq_len_cpu += 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( From a56c79bf79d5f8a067950c2c107761c1970c908c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 22 Jun 2026 14:34:27 +0000 Subject: [PATCH 037/214] delete --- lightllm/models/deepseek_v4/model.py | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index c824b24387..e435d64d3b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -33,7 +33,6 @@ from lightllm.distributed.communication_op import dist_group_manager logger = init_logger(__name__) -DSV4_DECODE_CUDAGRAPH_MAX_LEN = 8192 @ModelRegistry("deepseek_v4") @@ -100,17 +99,6 @@ def _init_mem_manager(self): self.req_manager.mem_manager = self.mem_manager return - def _init_cudagraph(self): - if not self.disable_cudagraph and self.graph_max_len_in_batch > DSV4_DECODE_CUDAGRAPH_MAX_LEN: - logger.info( - "DeepSeek-V4 caps decode cudagraph max_len_in_batch from %s to %s for the current " - "graph-safe sparse-attention path; longer decode batches run eager.", - self.graph_max_len_in_batch, - DSV4_DECODE_CUDAGRAPH_MAX_LEN, - ) - self.graph_max_len_in_batch = DSV4_DECODE_CUDAGRAPH_MAX_LEN - return super()._init_cudagraph() - def _init_att_backend(self): args = get_env_start_args() if args.llm_kv_type == "None": From e2869432b1fe502714182afa45e3b97acca0942d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 23 Jun 2026 03:02:40 +0000 Subject: [PATCH 038/214] item1: wire fused_q_indexer_rope_hadamard_quant (rope+hadamard+fp8quant+scale-fold in 1 kernel) 16k 320.29->352.14 (+9.9%, TTFT -14.6%), 64k 84.48->86.33 (+2.2%). Needle@16.5k PASS. --- .../layer_infer/transformer_layer_infer.py | 18 +++++++++--------- lightllm/models/deepseek_v4/model.py | 3 +++ 2 files changed, 12 insertions(+), 9 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 57b171a8d1..dc7454673c 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -547,22 +547,22 @@ def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, la Hadamard -> per-token fp8 quant. Returns (idx_q_fp8 [T,H,d], weights [T,H]); the per-token q fp8 scale and the head_dim^-0.5 * n_heads^-0.5 score scale are folded into weights -- the deep_gemm.fp8_mqa_logits contract (fp8 q carries no companion scale). Replicated -> full heads.""" - from lightllm.models.deepseek3_2.triton_kernel.act_quant import act_quant - from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform + # Fused: wq_b mm -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant, with the + # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: + # rotary_emb_fwd + hadamard_transform + act_quant + weights mul). freqs_cis is the compress rope + # table (same one the main compress-layer Q path uses); positions indexed inside the kernel. + from lightllm.third_party.sglang_jit.dsv4.elementwise import fused_q_indexer_rope_hadamard_quant - cos_tok = infer_state.position_cos_compress - sin_tok = infer_state.position_sin_compress token_num = q_lora.shape[0] if x.shape[0] != token_num: raise RuntimeError( f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" ) idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(token_num, self.index_n_heads, self.index_head_dim) - rotary_emb_fwd(idx_q[..., -self.qk_rope_head_dim :], None, cos_tok, sin_tok) - idx_q = hadamard_transform(idx_q, scale=self.index_head_dim ** -0.5) - idx_q_fp8, q_scale = act_quant(idx_q, self.index_head_dim, None) # fp8 [T,H,d], scale [T,H,1] - weights = layer_weight.idx_weights_proj_.mm(x).float() * self.indexer_weight_scale # [T, H] - weights = weights.unsqueeze(-1) * q_scale # fold per-token q scale + raw_w = layer_weight.idx_weights_proj_.mm(x).view(token_num, self.index_n_heads) # [T, H] raw + idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( + idx_q, raw_w, self.indexer_weight_scale, self.freqs_cis, infer_state.position_ids + ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded return idx_q_fp8, weights.squeeze(-1).contiguous() def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index e435d64d3b..af3ba3fc50 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -168,6 +168,9 @@ def _init_to_get_rotary(self): layer.freqs_cis = self._freqs_cis_compress if layer.compress_ratio else self._freqs_cis_sliding layer.cos_compress_table = self._cos_cached_compress layer.sin_compress_table = self._sin_cached_compress + # the indexer-Q fused kernel (compress rope) needs the complex compress freqs table. + if getattr(layer, "index_infer", None) is not None: + layer.index_infer.freqs_cis = self._freqs_cis_compress return From 58b145b7a71c471d0b8c4a0a06ca63aba4f9bd60 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 23 Jun 2026 03:46:29 +0000 Subject: [PATCH 039/214] item3: lazy-cache layer-independent c4 paged metadata (page_table/ctx_lens/meta/topk_lengths) across c4 layers needle@16.5k PASS; speed neutral (352.1->352.8 16k, 86.3->86.4 64k) but cuts ~20x metadata launches --- lightllm/models/deepseek_v4/infer_struct.py | 5 + .../layer_infer/transformer_layer_infer.py | 146 ++++++------------ .../triton_kernel/gather_c4_indexer_k_dsv4.py | 135 +--------------- 3 files changed, 58 insertions(+), 228 deletions(-) diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index ca2ac83b03..a6fe9339e0 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -23,8 +23,13 @@ def __init__(self): self.dsv4_sparse_req_idx = None self.dsv4_swa_indices = None self.dsv4_swa_lengths = None + # lazily-built (first c4 layer) cache of layer-independent paged-c4 metadata; reused by the + # other c4 layers in the same forward. Plain tuple (not a tensor attr) so copy_for_cuda_graph + # ignores it -- it's a capture-time wiring of layer0->others, not a staged graph input. + self._c4_paged_meta = None def init_some_extra_state(self, model): + self._c4_paged_meta = None # reset per forward before any c4 layer runs super().init_some_extra_state(model) # sets position_ids, b_q_seq_len, b_q_start_loc (prefill) pos = self.position_ids self.position_cos_sliding = torch.index_select(model._cos_cached_sliding, 0, pos) # [T, rope_dim//2] diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index dc7454673c..4266df4ae7 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,4 +1,3 @@ -import os import torch import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl @@ -15,6 +14,8 @@ from .compressor import prepare_compress_states from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo +import deep_gemm +from lightllm.third_party.sglang_jit.dsv4 import topk_transform_512 class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): @@ -566,10 +567,9 @@ def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, la return idx_q_fp8, weights.squeeze(-1).contiguous() def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): - """c4 scorer via ds3.2-style gather + deep_gemm.fp8_mqa_logits. Gather each request's causal c4 - keys into a padded-per-request ragged fp8 buffer (k row r*c4_cap+e), score every query token - over its absolute [ks, ke) range, then masked topk-512 -> c4 slots. Fixed shapes (c4_cap pinned - per graph bucket) keep the decode cuda graph capturable.""" + """c4 scorer via the page-safe deep_gemm.fp8_paged_mqa_logits over the paged c4 indexer pool, + then masked topk-512 -> c4 slots. Fixed shapes (c4_cap pinned per graph bucket) keep the decode + cuda graph capturable.""" mem_manager = infer_state.mem_manager index_topk = self.index_topk max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) @@ -591,105 +591,60 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, ) return slots.unsqueeze(1), lengths - b_req_idx = infer_state.b_req_idx - batch = b_req_idx.shape[0] - device = positions.device c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req - if os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER", "0") == "1": - out = self._c4_indices_paged( - infer_state=infer_state, - idx_q_fp8=idx_q_fp8, - weights=weights, - positions=positions, - c4_len=c4_len, - c4_cap=c4_cap, - ) - if out is not None: - return out - - import deep_gemm - from ..triton_kernel.gather_c4_indexer_k_dsv4 import gather_c4_indexer_k_ragged - - k_fp8, k_scale, ragged_slots = gather_c4_indexer_k_ragged( - mem_manager, - self.layer_idx_, - b_req_idx, - c4_len, - c4_cap, - infer_state.req_manager.req_to_token_indexs, - ) - # batch position of each query token -> absolute [ks, ke) into the padded buffer. - if infer_state.is_prefill: - token_batch_pos = torch.repeat_interleave( - torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len - ) - else: - token_batch_pos = torch.arange(batch, device=device, dtype=torch.int32) - valid_len = ((positions + 1) // 4).to(torch.int32) # causal candidate count per query - ks = token_batch_pos * c4_cap - ke = ks + valid_len - logits = deep_gemm.fp8_mqa_logits( - idx_q_fp8, (k_fp8, k_scale), weights, ks, ke, clean_logits=False, max_seqlen_k=c4_cap - ) # [T, c4_cap] f32, left-aligned: logits[t, j] = q_t . k[ks[t]+j] - col = torch.arange(c4_cap, device=device) - logits = logits.masked_fill(col.unsqueeze(0) >= valid_len.unsqueeze(1), float("-inf")) - top = logits.topk(index_topk, dim=-1).indices.to(torch.int32) # relative positions in [0, valid_len) - abs_idx = top + ks.unsqueeze(1) # absolute compact row - top_slots = ragged_slots[abs_idx.long()] # compact row -> c4 pool slot - invalid = top >= valid_len.unsqueeze(1) # topk over -inf padding when valid_len < index_topk - top_slots = torch.where(invalid, torch.full_like(top_slots, -1), top_slots) - topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) - return top_slots.unsqueeze(1), topk_lengths.contiguous() - - def _c4_indices_paged(self, infer_state, idx_q_fp8, weights, positions, c4_len, c4_cap): - import deep_gemm - from lightllm.third_party.sglang_jit.dsv4 import topk_transform_512 - from ..triton_kernel.gather_c4_indexer_k_dsv4 import build_c4_indexer_page_table - - mem_manager = infer_state.mem_manager - index_topk = self.index_topk device = positions.device - b_req_idx = infer_state.b_req_idx - batch = b_req_idx.shape[0] - validate = os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER_VALIDATE", "0") == "1" - validate_now = validate and not torch.cuda.is_current_stream_capturing() - - page_table, valid_flag = build_c4_indexer_page_table( - mem_manager, - b_req_idx, - c4_len, - c4_cap, - infer_state.req_manager.req_to_token_indexs, - infer_state.req_manager.HOLD_REQUEST_ID, - validate=validate_now, - ) - if validate_now and int(valid_flag.item()) == 0: - if os.getenv("LIGHTLLM_DSV4_PAGED_INDEXER_STRICT", "0") == "1": - raise RuntimeError("DeepSeek-V4 paged indexer requires page-aligned c4 slots") - return None + page_size = mem_manager.c4_indexer_pool.page_size + + # The page table / row_page_table / valid_len / ctx_lens / paged-logits metadata / topk_lengths + # are LAYER-INDEPENDENT (depend on request layout + c4_cap, not on weights/layer). Build them on + # the first c4 layer of the forward and reuse on the other ~20 c4 layers (was rebuilt per layer: + # build_c4_indexer_page_table + a [T,npages] gather + clamp/reshape + get_paged_mqa_logits_metadata + # each, i.e. ~20x redundant index/copy/clamp launches). Lazy (not init_some_extra_state) so it is + # computed inside the decode cuda graph with the capture-forced shapes -> no graph-cap mismatch. + cached = getattr(infer_state, "_c4_paged_meta", None) + if cached is None: + from ..triton_kernel.gather_c4_indexer_k_dsv4 import build_c4_indexer_page_table + + b_req_idx = infer_state.b_req_idx + batch = b_req_idx.shape[0] + page_table = build_c4_indexer_page_table( + mem_manager, + b_req_idx, + c4_len, + c4_cap, + infer_state.req_manager.req_to_token_indexs, + infer_state.req_manager.HOLD_REQUEST_ID, + ) - if infer_state.is_prefill: - token_batch_pos = torch.repeat_interleave( - torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len + if infer_state.is_prefill: + token_batch_pos = torch.repeat_interleave( + torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len + ) + row_page_table = page_table[token_batch_pos.long()].contiguous() + else: + row_page_table = page_table + + valid_len = ((positions + 1) // 4).to(torch.int32) + ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1).contiguous() + metadata = deep_gemm.get_paged_mqa_logits_metadata( + ctx_lens, + page_size, + deep_gemm.get_num_sms(), ) - row_page_table = page_table[token_batch_pos.long()].contiguous() - else: - row_page_table = page_table + topk_lengths = torch.clamp( + torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1 + ).contiguous() + cached = (row_page_table, valid_len, ctx_lens, metadata, topk_lengths) + infer_state._c4_paged_meta = cached - valid_len = ((positions + 1) // 4).to(torch.int32) - ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1).contiguous() + row_page_table, valid_len, ctx_lens, metadata, topk_lengths = cached kv_cache = mem_manager.c4_indexer_pool.get_layer_buffer(mem_manager.layer_to_c4_idx[self.layer_idx_]).view( mem_manager.c4_indexer_pool.num_pages, - mem_manager.c4_indexer_pool.page_size, + page_size, 1, self.index_head_dim + 4, ) - metadata = deep_gemm.get_paged_mqa_logits_metadata( - ctx_lens, - mem_manager.c4_indexer_pool.page_size, - deep_gemm.get_num_sms(), - ) logits = deep_gemm.fp8_paged_mqa_logits( idx_q_fp8.unsqueeze(1), kv_cache, @@ -706,7 +661,6 @@ def _c4_indices_paged(self, infer_state, idx_q_fp8, weights, positions, c4_len, valid_len, row_page_table, top_slots, - mem_manager.c4_indexer_pool.page_size, + page_size, ) - topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) - return top_slots.unsqueeze(1), topk_lengths.contiguous() + return top_slots.unsqueeze(1), topk_lengths diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py index a7a0a4be85..b50eab6726 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -3,112 +3,6 @@ import triton.language as tl -@triton.jit -def _gather_c4_indexer_k_kernel( - req_idx_ptr, # [batch] int — req_manager slot per batch position - c4_len_ptr, # [batch] int — number of causal c4 entries per request (= seq_len // ratio) - req_to_token_ptr, - req_to_token_stride0, - full_to_c4_ptr, - SlabFp8_ptr, # c4 indexer pool, viewed as fp8 (flat) - SlabF32_ptr, # same pool, viewed as f32 (flat) - Kout_fp8_ptr, # [batch*c4_cap, HEAD_DIM] fp8 - Kout_scale_ptr, # [batch*c4_cap] f32 - Slots_out_ptr, # [batch*c4_cap] int32 (compact->c4-slot map; -1 for padding) - c4_cap, - RATIO: tl.constexpr, - HEAD_DIM: tl.constexpr, - PAGE_SIZE: tl.constexpr, - BYTES_PER_PAGE: tl.constexpr, - SCALE_OFFSET: tl.constexpr, # page_size * head_dim (byte offset of the scale tail) -): - # entry index on grid-X (limit ~2^31), batch on grid-Y (<= running_max_req_size): c4_cap reaches - # 65536 at 256K context, which would blow the 65535 grid-Y cap if entries were the grid-Y axis. - e = tl.program_id(0) - r = tl.program_id(1) - # int64: out_pos*HEAD_DIM can exceed int32 at high batch + long context (read side is already int64). - out_pos = r.to(tl.int64) * c4_cap + e - c4_len = tl.load(c4_len_ptr + r) - if e >= c4_len: - # padding entry: mark slot invalid; K is never read (ke bounds the scorer range). - tl.store(Slots_out_ptr + out_pos, -1) - return - - # group-end token of compressed entry e lives at position e*RATIO + (RATIO-1). - req = tl.load(req_idx_ptr + r).to(tl.int64) - end_tok = e * RATIO + (RATIO - 1) - full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + end_tok).to(tl.int64) - c4_slot = tl.load(full_to_c4_ptr + full_slot).to(tl.int64) - valid = c4_slot >= 0 - - # inline PackedPagePool byte addressing (matches destindex_copy_indexer_k_dsv4 / gather_indexer_k): - # fp8 K at page*bytes_per_page + tok*head_dim; fp32 scale at (page*bytes_per_page + scale_off)//4 + tok. - page = c4_slot // PAGE_SIZE - tok = c4_slot % PAGE_SIZE - data_base = page * BYTES_PER_PAGE + tok * HEAD_DIM - scale_base = (page * BYTES_PER_PAGE + SCALE_OFFSET) // 4 + tok - - offs_d = tl.arange(0, HEAD_DIM) - k_fp8 = tl.load(SlabFp8_ptr + data_base + offs_d, mask=valid, other=0.0) - k_scale = tl.load(SlabF32_ptr + scale_base, mask=valid, other=0.0) - tl.store(Kout_fp8_ptr + out_pos * HEAD_DIM + offs_d, k_fp8) - tl.store(Kout_scale_ptr + out_pos, k_scale) - tl.store(Slots_out_ptr + out_pos, tl.where(valid, c4_slot, -1).to(tl.int32)) - - -@torch.no_grad() -def gather_c4_indexer_k_ragged( - mem_manager, - layer_index: int, - b_req_idx: torch.Tensor, - c4_len: torch.Tensor, - c4_cap: int, - req_to_token_indexs: torch.Tensor, -): - """Gather each request's causal c4 indexer keys into a padded-per-request ragged buffer for the - deep_gemm fp8_mqa_logits scorer (mirrors deepseek3_2's extract_indexer_ks, but reads our - PackedPagePool by c4 slot instead of a token-indexed [N,1,132] buffer). - - For batch position r and compressed entry e in [0, c4_len[r]): - c4_slot = full_to_c4[req_to_token[b_req_idx[r], e*ratio + (ratio-1)]] - The raw fp8 key + f32 scale at that slot land at row r*c4_cap + e of the output (so query token t - of request r reads keys [r*c4_cap, r*c4_cap + (pos+1)//ratio) -- absolute ks/ke offsets the caller - builds). Returns (k_fp8 [batch*c4_cap, HEAD_DIM] fp8, k_scale [batch*c4_cap] f32, slots - [batch*c4_cap] int32 = compact-row -> c4 pool slot, -1 for padding). Fixed shapes -> cuda-graph - safe (c4_cap is pinned per graph bucket); the padding region is never read by the scorer. - """ - pool = mem_manager.c4_indexer_pool - head_dim = mem_manager.indexer_head_dim - buf = pool.get_layer_buffer(mem_manager.layer_to_c4_idx[layer_index]).view(-1) - slab_fp8 = buf.view(torch.float8_e4m3fn) - slab_f32 = buf.view(torch.float32) - batch = b_req_idx.shape[0] - n = batch * c4_cap - k_fp8 = torch.empty((n, head_dim), dtype=torch.float8_e4m3fn, device=buf.device) - k_scale = torch.empty((n,), dtype=torch.float32, device=buf.device) - slots = torch.empty((n,), dtype=torch.int32, device=buf.device) - _gather_c4_indexer_k_kernel[(c4_cap, batch)]( - b_req_idx, - c4_len, - req_to_token_indexs, - req_to_token_indexs.stride(0), - mem_manager.full_to_c4_indexs, - slab_fp8, - slab_f32, - k_fp8, - k_scale, - slots, - c4_cap, - RATIO=4, - HEAD_DIM=head_dim, - PAGE_SIZE=pool.page_size, - BYTES_PER_PAGE=pool.bytes_per_page, - SCALE_OFFSET=pool.scale_offset_in_page, - num_warps=1, - ) - return k_fp8, k_scale, slots - - @triton.jit def _build_c4_indexer_page_table_kernel( req_idx_ptr, # [batch] int @@ -117,12 +11,10 @@ def _build_c4_indexer_page_table_kernel( req_to_token_stride0, full_to_c4_ptr, page_table_ptr, # [batch, page_cap] int32 - valid_flag_ptr, # [1] int32, initialized to 1; set to 0 on layout mismatch page_cap, hold_req_id, RATIO: tl.constexpr, PAGE_SIZE: tl.constexpr, - VALIDATE: tl.constexpr, ): p = tl.program_id(0) r = tl.program_id(1) @@ -141,22 +33,6 @@ def _build_c4_indexer_page_table_kernel( phys_page = c4_slot0 // PAGE_SIZE tl.store(page_table_ptr + r * page_cap + p, tl.where(active, phys_page, 0).to(tl.int32)) - if VALIDATE: - offs = tl.arange(0, PAGE_SIZE) - e = page_start + offs - valid = active & (e < c4_len) - full_pos = e * RATIO + (RATIO - 1) - full_slot = tl.load( - req_to_token_ptr + req * req_to_token_stride0 + full_pos, - mask=valid, - other=0, - ).to(tl.int64) - c4_slot = tl.load(full_to_c4_ptr + full_slot, mask=valid, other=-1).to(tl.int64) - expected = phys_page * PAGE_SIZE + offs - ok = tl.where(valid, (c4_slot == expected) & (c4_slot >= 0), True) - if tl.min(ok.to(tl.int32), axis=0) == 0: - tl.store(valid_flag_ptr, 0) - @torch.no_grad() def build_c4_indexer_page_table( @@ -166,14 +42,12 @@ def build_c4_indexer_page_table( c4_cap: int, req_to_token_indexs: torch.Tensor, hold_req_id: int, - validate: bool = False, ): """Build the logical-c4-page -> physical-c4-page table expected by DeepGEMM paged logits. - This is safe only when each logical c4 page maps to a physical page with matching offsets: + Safe only when each logical c4 page maps to a physical page with matching offsets: c4_slot(entry p*64 + o) == page_table[p] * 64 + o - The optional validation flag checks that invariant and lets the caller fall back to the - gather path while we keep the current token-slot allocator. + which the current token-slot allocator guarantees. """ pool = mem_manager.c4_indexer_pool page_size = pool.page_size @@ -181,7 +55,6 @@ def build_c4_indexer_page_table( batch = b_req_idx.shape[0] page_cap = c4_cap // page_size page_table = torch.empty((batch, page_cap), dtype=torch.int32, device=b_req_idx.device) - valid_flag = torch.ones((1,), dtype=torch.int32, device=b_req_idx.device) _build_c4_indexer_page_table_kernel[(page_cap, batch)]( b_req_idx, c4_len, @@ -189,12 +62,10 @@ def build_c4_indexer_page_table( req_to_token_indexs.stride(0), mem_manager.full_to_c4_indexs, page_table, - valid_flag, page_cap, int(hold_req_id), RATIO=4, PAGE_SIZE=page_size, - VALIDATE=validate, num_warps=1, ) - return page_table, valid_flag + return page_table From a0379bbad2dbe1407ffcdd28a3e97a0dfa653ae6 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 23 Jun 2026 05:02:22 +0000 Subject: [PATCH 040/214] gate-bf16 (flag) + drop redundant attn_sink fp32 copy + lazy gen_nsa_ks_ke 16k 352.8->368.1 (+4.3%), 64k 86.4->90.0 (+4.2%). needle@16.5k PASS; gsm8k 8-shot 300q acc=0.900/0-invalid (== fp32 baseline ~0. 902) with LIGHTLLM_DSV4_GATE_BF16=1. --- .../attention/nsa/fp8_flashmla_sparse.py | 18 ++++++++++++++++-- .../layer_infer/transformer_layer_infer.py | 1 + .../layer_infer/transformer_layer_infer.py | 2 +- .../layer_weights/transformer_layer_weight.py | 7 +++---- 4 files changed, 21 insertions(+), 7 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index dc18ecf4ba..5a3f125425 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -88,6 +88,15 @@ class NsaFlashMlaFp8SparsePrefillAttState(BasePrefillAttState): def init_state(self): self.backend: NsaFlashMlaFp8SparseAttBackend = self.backend + return + + def ensure_nsa_ks_ke(self): + """Build the ragged ks/ke/lengths (+ ragged_mem_index) the DeepSeek-3.2 indexer consumes. The + indexer calls this explicitly before reading them; DeepSeek-V4 uses its own indexer and never + calls it, so V4 prefill skips the alloc + gen_nsa_ks_ke kernel. Idempotent + layer-independent: + the first call in a forward computes, the other layers reuse.""" + if self.ks is not None: + return self.ragged_mem_index = torch.empty( self.infer_state.total_token_num, dtype=torch.int32, @@ -169,7 +178,7 @@ def _nsa_prefill_att( return mla_out def _flashmla_kvcache_prefill_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: - attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() + attn_sink = nsa_dict["attn_sink"] metadata = _metadata_from_dict(self.infer_state, nsa_dict) return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) @@ -252,6 +261,11 @@ def init_state(self): self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} return + def ensure_nsa_ks_ke(self): + # decode builds ks/ke eagerly in init_state (outside the cuda graph, for capture safety), so + # they are already available -- this satisfies the shared DeepSeek-3.2 indexer ensure contract. + return + def reset_sched_meta_for_capture(self): # cuda-graph capture hook: the warmup pass already locked/stored sched meta on this # (shared) state object; reset so the capture pass re-plans INSIDE the graph and every @@ -329,7 +343,7 @@ def _nsa_decode_att( return o_tensor[:, 0, :, :] # [b, 1, h, d] -> [b, h, d] def _flashmla_kvcache_decode_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: - attn_sink = nsa_dict["attn_sink"].to(torch.float32).contiguous() + attn_sink = nsa_dict["attn_sink"] metadata = _metadata_from_dict(self.infer_state, nsa_dict) return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) diff --git a/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py index d6eaebe2fd..8c7506b1e1 100644 --- a/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py @@ -206,6 +206,7 @@ def _get_indices( weights = layer_weight.weights_proj_.mm(hidden_states) * self.index_n_heads_scale weights = weights.unsqueeze(-1) * q_scale + att_state.ensure_nsa_ks_ke() ks = att_state.ks ke = att_state.ke lengths = att_state.lengths diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 4266df4ae7..f5b2ad5463 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -307,7 +307,7 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV if not self.enable_ep_moe: x = self._tpsp_allgather(input=x, infer_state=infer_state) - logits = layer_weight.gate_weight_.mm(x.float()).contiguous() + logits = layer_weight.gate_weight_.mm(x).float().contiguous() weights, indices = self._select_experts(logits, infer_state, layer_weight) # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 5896027b38..e818351c2e 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -189,14 +189,13 @@ def _init_indexer(self): # ------------------------------------------------------------------ moe def _init_moe(self): p = f"{self.prefix}.ffn" - # router gate (replicated). Stored as fp32: the topk_hash_softplus_sqrt router wants fp32 logits, - # so keep the gate matmul in fp32 — but store the (constant) weight as fp32 once here instead of - # re-casting it to fp32 on every forward in _ffn. + # Router gate in bf16 (matches the sglang/vLLM DeepSeek references, which run the gate GEMM in + # the model dtype); the bf16 GEMM output is cast back to fp32 in _ffn for topk_hash_softplus_sqrt. self.gate_weight_ = ROWMMWeight( in_dim=self.hidden, out_dims=[self.n_routed_experts], weight_names=f"{p}.gate.weight", - data_type=torch.float32, + data_type=torch.bfloat16, quant_method=None, tp_rank=0, tp_world_size=1, From b796d48da6f0a3496057e7ba37f4b0647e8448e8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 23 Jun 2026 05:07:01 +0000 Subject: [PATCH 041/214] cache prefill FlashMLA sched-meta per compress-ratio (was rebuilt every attention layer) 16k 368.1->416.4 (+13.1%, TTFT -17.5%), 64k 90.0->91.4. needle ok; gsm8k 8-shot 300q acc=0.900/0-invalid. The per-layer get_mla_ metadata (~435ms) was a host-side planning bubble x42 layers. --- .../basemodel/attention/nsa/fp8_flashmla_sparse.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 5a3f125425..115d5d1934 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -85,9 +85,11 @@ class NsaFlashMlaFp8SparsePrefillAttState(BasePrefillAttState): ke: torch.Tensor = None lengths: torch.Tensor = None ragged_mem_index: torch.Tensor = None + flashmla_sched_meta: object = None def init_state(self): self.backend: NsaFlashMlaFp8SparseAttBackend = self.backend + self.flashmla_sched_meta = {} return def ensure_nsa_ks_ke(self): @@ -115,6 +117,15 @@ def ensure_nsa_ks_ke(self): ) return + def _get_flashmla_sched_meta(self, compress_ratio: int): + sched_meta = self.flashmla_sched_meta.get(compress_ratio) + if sched_meta is None: + import flash_mla + + sched_meta = flash_mla.get_mla_metadata()[0] + self.flashmla_sched_meta[compress_ratio] = sched_meta + return sched_meta + def prefill_att( self, q: torch.Tensor, @@ -196,7 +207,7 @@ def _flashmla_kvcache_att( q_4d = q.unsqueeze(1).contiguous() q_4d, attn_sink, num_real_heads = _pad_q_heads(q_4d, attn_sink) k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) - sched_meta, _ = flash_mla.get_mla_metadata() + sched_meta = self._get_flashmla_sched_meta(nsa_dict["compress_ratio"]) out, _ = flash_mla.flash_mla_with_kvcache( q=q_4d, k_cache=k_cache, From e07b85e59b2f2d9beb5d2be06ae934f62a29932c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 23 Jun 2026 09:39:37 +0000 Subject: [PATCH 042/214] 2-stream --- lightllm/models/deepseek_v4/infer_struct.py | 8 ++ .../deepseek_v4/layer_infer/compressor.py | 4 +- .../layer_infer/transformer_layer_infer.py | 91 +++++++++++++++---- lightllm/models/deepseek_v4/model.py | 4 + 4 files changed, 86 insertions(+), 21 deletions(-) diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index a6fe9339e0..f37620ddff 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -23,6 +23,8 @@ def __init__(self): self.dsv4_sparse_req_idx = None self.dsv4_swa_indices = None self.dsv4_swa_lengths = None + # token -> batch-position map for the compressor; built per prefill forward in init_some_extra_state. + self._dsv4_token_to_batch_idx = None # lazily-built (first c4 layer) cache of layer-independent paged-c4 metadata; reused by the # other c4 layers in the same forward. Plain tuple (not a tensor attr) so copy_for_cuda_graph # ignores it -- it's a capture-time wiring of layer0->others, not a staged graph input. @@ -40,8 +42,14 @@ def init_some_extra_state(self, model): # Layer-independent; the swa kernel + build_metadata's c4/c128 readers all reuse it. if self.is_prefill: self.dsv4_sparse_req_idx = torch.repeat_interleave(self.b_req_idx, self.b_q_seq_len.long()) + self._dsv4_token_to_batch_idx = torch.repeat_interleave( + torch.arange(self.b_req_idx.shape[0], device=self.b_req_idx.device), + self.b_q_seq_len.long(), + output_size=pos.numel(), + ).to(torch.int32) else: self.dsv4_sparse_req_idx = self.b_req_idx + self._dsv4_token_to_batch_idx = None # Sliding-window indices: layer-independent (full_to_swa is global, window is const), so build # once here via one fused kernel instead of recomputing per layer. const [T, window] shape is # cuda-graph-safe (no max_kv_seq_len dependence) and auto-staged by copy_for_cuda_graph. diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index c66fd22de5..695ac33bc8 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -330,7 +330,9 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, if token_to_batch_idx is None or token_to_batch_idx.numel() != infer_state.position_ids.numel(): q_lens = (infer_state.b_seq_len - infer_state.b_ready_cache_len).to(torch.long) batch_idx = torch.arange(infer_state.b_req_idx.shape[0], device=infer_state.b_req_idx.device) - token_to_batch_idx = torch.repeat_interleave(batch_idx, q_lens).to(torch.int32) + token_to_batch_idx = torch.repeat_interleave( + batch_idx, q_lens, output_size=infer_state.position_ids.numel() + ).to(torch.int32) infer_state._dsv4_token_to_batch_idx = token_to_batch_idx return CoreCompressorMetadata( diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index f5b2ad5463..464639f0f7 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -55,6 +55,7 @@ def __init__(self, layer_num, network_config): self.index_infer = DeepseekV4IndexInfer( layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ ) + self.dsv4_prefill_aux_stream = None # ------------------------------------------------------------------ forward (HC-threaded) def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): @@ -222,19 +223,44 @@ def att_func(new_infer_state: DeepseekV4InferStateInfo): return self._context_attention_kernel(q, q_lora, x, infer_state, layer_weight) + def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight): + cos_table, sin_table = self.cos_compress_table, self.sin_compress_table + aux_stream = self.dsv4_prefill_aux_stream + if self.compress_ratio == 4 and aux_stream is not None and not torch.cuda.is_current_stream_capturing(): + # _dsv4_token_to_batch_idx is built in init_some_extra_state (default stream, before this fork), + # so both the aux indexer-compressor and the main compressor read a ready, race-free tensor. + main_stream = torch.cuda.current_stream() + aux_stream.wait_stream(main_stream) # fork: aux waits for x / q_lora produced on main + with torch.cuda.stream(aux_stream): + # x / q_lora are main-allocated and read here -> record so the allocator won't reuse them. + x.record_stream(aux_stream) + q_lora.record_stream(aux_stream) + self.index_infer.write_indexer_k( + x, infer_state, layer_weight, cos_table, sin_table, use_custom_tensor_manager=False + ) + meta = self.index_infer.build_metadata( + x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=False + ) + self.compressor.prepare_states(x, infer_state, layer_weight) + self.compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) + main_stream.wait_stream(aux_stream) # join before prefill_att reads the indices / latent KV + # extra_indices / extra_lengths were allocated on aux -> record on main so they survive until consumed. + for _t in (meta.get("extra_indices"), meta.get("extra_lengths")): + if _t is not None: + _t.record_stream(main_stream) + return meta + + # serial fallback -- semantics identical to the original sequence. + self.compressor.prepare_states(x, infer_state, layer_weight) + self.compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) + # write c4 Lightning-Indexer keys BEFORE build_metadata so the scorer reads fresh+accumulated entries. + self.index_infer.write_indexer_k(x, infer_state, layer_weight, cos_table, sin_table) + return self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) + def _context_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - self.compressor.prepare_states(x, infer_state, layer_weight) - self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) - # Write this step's c4 Lightning-Indexer keys (no-op off c4) BEFORE build_metadata so the - # scorer (gather + deep_gemm.fp8_mqa_logits) reads fresh+accumulated entries from the indexer pool. - self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) - # Build the FINAL flash_mla index tensors here (model side), so att_control is a thin - # transport of ready-to-forward tensors -- not indexer raw material. Must stay after - # fused_compress (c4 reads the indexer-K pool it writes) and before prefill_att (keeps the - # c4 scorer/topk at the same cuda-graph capture position). - meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) + meta = self._compress_and_index(q_lora, x, infer_state, layer_weight) att_control = AttControl( nsa_prefill=True, nsa_prefill_dict={ @@ -391,7 +417,10 @@ def prepare_states( x: torch.Tensor, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, + use_custom_tensor_manager: bool = True, ): + # use_custom_tensor_manager=False routes the .mm outputs through torch.empty (stream-aware) + # instead of the stream-blind global cache -- required when this runs on the prefill aux stream. self._metadata = prepare_compress_states( infer_state=infer_state, layer_idx=self.layer_idx_, @@ -402,8 +431,8 @@ def prepare_states( if self.is_in_indexer: # indexer wkv/wgate are two separate replicated weights; cat -> [T, 2*coff*idx_hd] # (same [kv | score] layout the fused compressor_wkv_gate_ produces for attention). - kv = layer_weight.idx_cmp_wkv_.mm(x) - gate = layer_weight.idx_cmp_wgate_.mm(x) + kv = layer_weight.idx_cmp_wkv_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager) + gate = layer_weight.idx_cmp_wgate_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager) self._metadata.kv_score = torch.cat([kv, gate], dim=-1).float() ape = layer_weight.idx_cmp_ape_.weight else: @@ -483,14 +512,24 @@ def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): else None ) - def write_indexer_k(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight, cos_table, sin_table): + def write_indexer_k( + self, + x, + infer_state: DeepseekV4InferStateInfo, + layer_weight, + cos_table, + sin_table, + use_custom_tensor_manager=True, + ): """c4-only: compress this step's tokens into per-c4-entry indexer keys and pack them into c4_indexer_pool. MUST run before build_metadata so the scorer (gather + deep_gemm.fp8_mqa_logits) reads the finished entries; runs every step (incl. in the decode graph) so keys accumulate for later long-context scoring. No-op on c128 / dense layers.""" if self.compress_ratio != 4: return - self.indexer_compressor.prepare_states(x, infer_state, layer_weight) + self.indexer_compressor.prepare_states( + x, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager + ) self.indexer_compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) scratch = self.indexer_compressor._metadata.out_buffer # [T, index_head_dim] bf16 (group-end rows valid) # Rotate K (post norm+rope) by the SAME 1/sqrt(d) Hadamard the q kernel applies, so @@ -507,7 +546,9 @@ def write_indexer_k(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight masked_slots = torch.where(completed, out_slots, torch.full_like(out_slots, -1)).to(torch.int32) mem_manager.pack_indexer_k_to_cache(self.layer_idx_, masked_slots, scratch) - def build_metadata(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + def build_metadata( + self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + ): """Return the final flash_mla index tensors for this layer's compress variant. swa indices and the per-token req_idx are layer-independent and precomputed once in init_some_extra_state (read here); only the c4 scorer / c128 gather is per-layer. The backend pairs these with the @@ -518,7 +559,9 @@ def build_metadata(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer positions = infer_state.position_ids extra_indices = extra_lengths = None if self.compress_ratio == 4: - idx_q_fp8, weights = self._indexer_q_weight(x, q_lora, infer_state, layer_weight) + idx_q_fp8, weights = self._indexer_q_weight( + x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager + ) extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: extra_indices, extra_lengths = self._c128_indices(infer_state, req_idx, positions) @@ -543,7 +586,9 @@ def _c128_indices(self, infer_state: DeepseekV4InferStateInfo, req_idx, position ) return indices.unsqueeze(1), lengths - def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight): + def _indexer_q_weight( + self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + ): """fp8 indexer q (mirrors deepseek3_2 NsaInfer): wq_b -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant. Returns (idx_q_fp8 [T,H,d], weights [T,H]); the per-token q fp8 scale and the head_dim^-0.5 * n_heads^-0.5 score scale are folded into weights -- the @@ -559,8 +604,12 @@ def _indexer_q_weight(self, x, q_lora, infer_state: DeepseekV4InferStateInfo, la raise RuntimeError( f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" ) - idx_q = layer_weight.idx_wq_b_.mm(q_lora).view(token_num, self.index_n_heads, self.index_head_dim) - raw_w = layer_weight.idx_weights_proj_.mm(x).view(token_num, self.index_n_heads) # [T, H] raw + idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads, self.index_head_dim + ) + raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads + ) # [T, H] raw idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( idx_q, raw_w, self.indexer_weight_scale, self.freqs_cis, infer_state.position_ids ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded @@ -619,7 +668,9 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, if infer_state.is_prefill: token_batch_pos = torch.repeat_interleave( - torch.arange(batch, device=device, dtype=torch.int32), infer_state.b_q_seq_len + torch.arange(batch, device=device, dtype=torch.int32), + infer_state.b_q_seq_len, + output_size=positions.numel(), ) row_page_table = page_table[token_batch_pos.long()].contiguous() else: diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index af3ba3fc50..a86da8774e 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -112,6 +112,10 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() + if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1": + prefill_aux_stream = torch.cuda.Stream() + for layer in self.layers_infer: + layer.dsv4_prefill_aux_stream = prefill_aux_stream dist_group_manager.new_deepep_group( self.config["n_routed_experts"], self.config["hidden_size"], From 51f5c8450281e023008bbad76f45de26ca3a82e5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 24 Jun 2026 02:11:09 +0000 Subject: [PATCH 043/214] fix parser --- lightllm/models/deepseek_v4/model.py | 5 ++++- lightllm/server/api_cli.py | 1 + lightllm/server/api_models.py | 2 +- lightllm/server/api_openai.py | 9 +++++++-- lightllm/server/build_prompt.py | 2 ++ lightllm/server/core/objs/start_args_type.py | 1 + lightllm/server/function_call_parser.py | 3 ++- lightllm/server/reasoning_parser.py | 1 + 8 files changed, 19 insertions(+), 5 deletions(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index a86da8774e..1be136167a 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -255,13 +255,16 @@ def apply_chat_template( if thinking is None: thinking = bool(enable_thinking) if enable_thinking is not None else False thinking_mode = "thinking" if thinking else "chat" + effort = kwargs.get("reasoning_effort") + if effort not in ("max", "high", None): + effort = None encoding = self._get_encoding_module() prompt = encoding.encode_messages( msgs, thinking_mode=thinking_mode, drop_thinking=kwargs.get("drop_thinking", True), add_default_bos_token=kwargs.get("add_default_bos_token", True), - reasoning_effort=kwargs.get("reasoning_effort"), + reasoning_effort=effort, ) if tokenize: diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e745492173..70a786edf6 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -175,6 +175,7 @@ def make_argument_parser() -> argparse.ArgumentParser: choices=[ "deepseek-r1", "deepseek-v3", + "deepseek-v4", "glm45", "gpt-oss", "kimi", diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index 1737d2774d..d36ad0fdde 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -221,7 +221,7 @@ class ChatCompletionRequest(BaseModel): parallel_tool_calls: Optional[bool] = True # OpenAI parameters for reasoning and others - reasoning_effort: Optional[Literal["low", "medium", "high"]] = None + reasoning_effort: Optional[Literal["low", "medium", "high", "max"]] = None chat_template_kwargs: Optional[Dict] = None separate_reasoning: Optional[bool] = True stream_reasoning: Optional[bool] = False diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index 0d934c44c9..1df04af3ab 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -165,8 +165,13 @@ def _is_force_thinking_mode(request: ChatCompletionRequest) -> bool: return False if reasoning_parser in ["qwen3-thinking", "gpt-oss", "minimax"]: return True - if reasoning_parser in ["deepseek-v3"]: - return request.chat_template_kwargs is not None and request.chat_template_kwargs.get("thinking") is True + if reasoning_parser in ["deepseek-v3", "deepseek-v4"]: + chat_template_kwargs = request.chat_template_kwargs or {} + if "thinking" in chat_template_kwargs: + return chat_template_kwargs["thinking"] is True + if request.reasoning_effort is not None: + return request.reasoning_effort != "none" + return False if reasoning_parser in ["qwen3", "glm45", "nano_v3", "interns1", "gemma4"]: # qwen3, glm45, nano_v3, interns1, and gemma4 are reasoning by default; return not request.chat_template_kwargs or request.chat_template_kwargs.get("enable_thinking", True) is True diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index 0565a8f0cf..bda1aafea5 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -151,6 +151,8 @@ async def build_prompt(request, tools) -> str: if request.chat_template_kwargs: kwargs.update(request.chat_template_kwargs) + if request.reasoning_effort is not None and "reasoning_effort" not in kwargs: + kwargs["reasoning_effort"] = request.reasoning_effort # 修复一些parser类型是默认打开thinking,但是 tokenizer有时候不知道打开了thinking。导致 # 构建的reasoning parser 和 tokenizer 的行为不对齐导致的问题。 from .api_openai import _is_force_thinking_mode diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 40c8028158..8f2df6eba8 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -42,6 +42,7 @@ class StartArgs: "choices": [ "deepseek-r1", "deepseek-v3", + "deepseek-v4", "glm45", "gpt-oss", "kimi", diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index 63c9f6ac8f..c9d06b18ec 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1698,9 +1698,10 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami if param_matches and len(param_matches) > len(self._accumulated_params): self._accumulated_params = param_matches current_args_json = self._dsml_params_to_json(param_matches) + open_args_json = current_args_json[:-1] # drop trailing '}' sent = len(self.streamed_args_for_tool[self.current_tool_id]) - argument_diff = current_args_json[sent:] + argument_diff = open_args_json[sent:] if argument_diff: calls.append( diff --git a/lightllm/server/reasoning_parser.py b/lightllm/server/reasoning_parser.py index 8a8d07355b..f351d8a6c8 100644 --- a/lightllm/server/reasoning_parser.py +++ b/lightllm/server/reasoning_parser.py @@ -903,6 +903,7 @@ class ReasoningParser: DetectorMap: Dict[str, Type[BaseReasoningFormatDetector]] = { "deepseek-r1": DeepSeekR1Detector, "deepseek-v3": Qwen3Detector, + "deepseek-v4": Qwen3Detector, "glm45": Qwen3Detector, "gpt-oss": GptOssDetector, "kimi": KimiDetector, From 2f12a077be2906f0a23ca58f4d5a0af0ae3afaf4 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 24 Jun 2026 02:32:14 +0000 Subject: [PATCH 044/214] fix multi-invoke --- lightllm/server/api_models.py | 2 +- lightllm/server/function_call_parser.py | 9 +++++---- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index d36ad0fdde..dcf3d2a0aa 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -221,7 +221,7 @@ class ChatCompletionRequest(BaseModel): parallel_tool_calls: Optional[bool] = True # OpenAI parameters for reasoning and others - reasoning_effort: Optional[Literal["low", "medium", "high", "max"]] = None + reasoning_effort: Optional[Literal["none", "low", "medium", "high", "max"]] = None chat_template_kwargs: Optional[Dict] = None separate_reasoning: Optional[bool] = True stream_reasoning: Optional[bool] = False diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index c9d06b18ec..c08cb83fa3 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1593,8 +1593,10 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami try: # Try to find complete invoke blocks first - complete_invoke_match = self.invoke_regex.search(current_text) - if complete_invoke_match: + while True: + complete_invoke_match = self.invoke_regex.search(current_text) + if not complete_invoke_match: + break func_name = complete_invoke_match.group(1) invoke_body = complete_invoke_match.group(2) @@ -1658,8 +1660,7 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami self.current_tool_name_sent = False self._accumulated_params = [] self.streamed_args_for_tool.append("") - - return StreamingParseResult(normal_text="", calls=calls) + current_text = self._buffer # Partial invoke: name is known but parameters are still streaming partial_match = self.partial_invoke_regex.search(current_text) From da3fec227590a978ea14cc2fcb062410243434c3 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 24 Jun 2026 10:10:14 +0000 Subject: [PATCH 045/214] speed up prepare --- lightllm/common/req_manager.py | 166 ++++++++++++++++++++++++--------- 1 file changed, 124 insertions(+), 42 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index e633f9ee3d..d40fe96989 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -422,6 +422,7 @@ def __init__( self.indexer_head_dim = indexer_head_dim self.layer_to_c4_idx = {} self.layer_to_c128_idx = {} + self.mem_manager = mem_manager c4 = c128 = 0 for lid, r in enumerate(self.compress_rates): if r == 4: @@ -457,7 +458,7 @@ def prepare_prefill_swa( 本 chunk 起点 L = ready_cache_len,首个新 token(位置 L)的窗口是 [L-W+1, L];回收 边界再额外保留一个 radix 页(_swa_retain_len),即位置 < L-retain+1。先回收再分配。 必须在 init_req_to_token_indexes 之后调用(位置对齐分配经 req_to_token 行派生/scatter)。""" - assert self.mem_manager is not None + self.mem_manager: DeepseekV4MemoryManager if self.sliding_window is not None: retain = self._swa_retain_len() evict_slots = [] @@ -575,6 +576,15 @@ def _compress_mapping_alloc(self, ratio: int): return self.mem_manager.full_to_c128_indexs, self.mem_manager.alloc_c128 raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") + def _c4_group_end_full_slots(self, req_rows, entries: torch.Tensor) -> torch.Tensor: + """组末 token 的 full 槽位 id (token 位置 = entry*4+3);req_rows 可为标量 req_idx 或行张量。""" + return self.req_to_token_indexs[req_rows, entries * 4 + 3].long() + + def _register_c4_slots(self, full_slots: torch.Tensor, slots: torch.Tensor) -> None: + """写入 full->c4 槽映射并按页累加存活计数。""" + self.mem_manager.full_to_c4_indexs[full_slots] = slots + self.mem_manager.count_c4_slots(slots, 1) + def _scatter_c4_prefill_slots_slow(self, req_idx: int, first: int, last: int) -> None: """Idempotence fallback for overlapped/repeated c4 prep.""" page = DSV4_C4_PAGE_SIZE @@ -583,7 +593,7 @@ def _scatter_c4_prefill_slots_slow(self, req_idx: int, first: int, last: int) -> e0 = max(first, page_base) e1 = min(last, page_base + page) entries = torch.arange(e0, e1, dtype=torch.long, device="cuda") - full_slots = self.req_to_token_indexs[req_idx, entries * 4 + 3].long() + full_slots = self._c4_group_end_full_slots(req_idx, entries) existing = mapping[full_slots] missing = existing < 0 if not bool(missing.any()): @@ -604,8 +614,7 @@ def _scatter_c4_prefill_slots_slow(self, req_idx: int, first: int, last: int) -> slots = (base + entries % page).to(torch.int32) if mapped.numel() > 0: assert bool((existing[existing >= 0] == slots[existing >= 0]).all()) - mapping[full_slots[missing]] = slots[missing] - self.mem_manager.count_c4_slots(slots[missing], 1) + self._register_c4_slots(full_slots[missing], slots[missing]) return def _scatter_c4_prefill_slots(self, req_idx: int, first: int, last: int) -> None: @@ -616,51 +625,132 @@ def _scatter_c4_prefill_slots(self, req_idx: int, first: int, last: int) -> None """ if last <= first: return - page = DSV4_C4_PAGE_SIZE mapping = self.mem_manager.full_to_c4_indexs - entries = torch.arange(first, last, dtype=torch.long, device="cuda") - full_slots = self.req_to_token_indexs[req_idx, entries * 4 + 3].long() + full_slots = self._c4_group_end_full_slots(req_idx, torch.arange(first, last, device="cuda")) need = mapping[full_slots] < 0 if not bool(need.any()): return if not bool(need.all()): self._scatter_c4_prefill_slots_slow(req_idx, first, last) return + self._scatter_c4_prefill_slots_fresh(req_idx, first, last) + return + def _scatter_c4_prefill_slots_fresh(self, req_idx: int, first: int, last: int) -> None: + """Sync-free fast path for entries [first, last) the caller already knows are all fresh: + page-safe alloc with the continuation base read on-GPU. KvCacheAllocator.alloc is CPU-side + (watermark + pinned buffer, no D2H), so this has no syncs.""" + page = DSV4_C4_PAGE_SIZE + mapping = self.mem_manager.full_to_c4_indexs + entries = torch.arange(first, last, dtype=torch.long, device="cuda") + full_slots = self._c4_group_end_full_slots(req_idx, entries) first_page = first // page - last_page = (last - 1) // page - n_pages = last_page - first_page + 1 + n_pages = (last - 1) // page - first_page + 1 bases = torch.empty((n_pages,), dtype=torch.long, device="cuda") - base_start = 0 - if first % page != 0: + if first % page != 0: # chunk starts mid-page -> continue the prev chunk's physical page prev_full = self.req_to_token_indexs[req_idx, first * 4 - 1].long() - prev_slot = int(mapping[prev_full].item()) - assert prev_slot >= 0 and prev_slot % page == (first - 1) % page - bases[0] = prev_slot - ((first - 1) % page) + bases[0] = mapping[prev_full].long() - ((first - 1) % page) base_start = 1 - - new_page_count = n_pages - base_start - if new_page_count > 0: - new_pages = self.mem_manager.alloc_c4_pages(new_page_count).cuda(non_blocking=True).long() - bases[base_start:] = new_pages * page - + if n_pages - base_start > 0: + bases[base_start:] = ( + self.mem_manager.alloc_c4_pages(n_pages - base_start).cuda(non_blocking=True).long() * page + ) page_local = torch.div(entries, page, rounding_mode="floor") - first_page slots = (bases[page_local] + entries % page).to(torch.int32) - mapping[full_slots] = slots - self.mem_manager.count_c4_slots(slots, 1) + self._register_c4_slots(full_slots, slots) return - def _scatter_c4_decode_slots( - self, - b_req_idx_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, - mem_indexes: torch.Tensor, - ) -> None: + def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list) -> None: + """Whole-batch c4 prefill scatter in O(1) GPU ops (independent of request count). The per-req + loop cost O(N) launches + 2-3 D2H syncs each; here every request's group-end entries are + flattened (ragged) and processed in one gather / idempotency-check / page-alloc / scatter / + count. Falls back to the per-req idempotent path on partial/re-run; preserves the page + invariant (logical entry e -> physical_page*64 + e%64, same logical page shares a physical + page). KvCacheAllocator.alloc is CPU-side so one batched alloc has no D2H.""" + page = DSV4_C4_PAGE_SIZE + mapping = self.mem_manager.full_to_c4_indexs + device = mapping.device + + # host plan (cheap int arithmetic, no GPU). dup req in one call would break vectorized + # continuation ordering -> fall back to the safe per-req path. + plan, seen, duplicate_req = [], set(), False + for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): + req_idx = int(req_idx) + if req_idx == self.HOLD_REQUEST_ID: + continue + first, last = int(ready_len) // 4, int(seq_len) // 4 + if last <= first: + continue + duplicate_req |= req_idx in seen + seen.add(req_idx) + plan.append((req_idx, first, last)) + if not plan: + return + if duplicate_req: + for req_idx, first, last in plan: + self._scatter_c4_prefill_slots(req_idx, first, last) + return + + def to_cuda_long(key, data): + return g_pin_mem_manager.gen_from_list(key=key, data=data, dtype=torch.int64).to(device, non_blocking=True) + + reqs, firsts, lasts = zip(*plan) + counts = [last - first for first, last in zip(firsts, lasts)] + first_pages = [first // page for first in firsts] + page_counts = [((last - 1) // page) - fp + 1 for last, fp in zip(lasts, first_pages)] + page_offsets, total_pages = [], 0 + for n_pages in page_counts: + page_offsets.append(total_pages) + total_pages += n_pages + total_entries = sum(counts) + + # one pinned H2D copy for all per-request metadata (5 cols), then per-entry ragged expansion + meta = to_cuda_long( + "dsv4_c4_prefill_meta", + [x for row in zip(reqs, firsts, first_pages, counts, page_offsets) for x in row], + ).view(-1, 5) + reqs_t, firsts_t, first_pages_t, counts_t, page_offsets_t = meta.unbind(1) + seg = torch.repeat_interleave(torch.arange(len(plan), device=device), counts_t, output_size=total_entries) + seg_starts = counts_t.cumsum(0) - counts_t + entries = firsts_t[seg] + torch.arange(total_entries, device=device) - seg_starts[seg] + full_slots = self._c4_group_end_full_slots(reqs_t[seg], entries) + + if not bool((mapping[full_slots] < 0).all()): # the single batched idempotency sync + for req_idx, first, last in plan: + self._scatter_c4_prefill_slots(req_idx, first, last) + return + + # physical base per logical page: fresh pages from one alloc; mid-page continuations read prev + cont = [(off, req, first) for off, req, first in zip(page_offsets, reqs, firsts) if first % page != 0] + if not cont: + page_bases = self.mem_manager.alloc_c4_pages(total_pages).to(device, non_blocking=True).long() * page + else: + page_bases = torch.empty(total_pages, dtype=torch.long, device=device) + new_pos = [ + pos + for off, n_pages, first in zip(page_offsets, page_counts, firsts) + for pos in range(off + int(first % page != 0), off + n_pages) + ] + if new_pos: + new_pos_t = to_cuda_long("dsv4_c4_prefill_new_pos", new_pos) + page_bases[new_pos_t] = ( + self.mem_manager.alloc_c4_pages(len(new_pos)).to(device, non_blocking=True).long() * page + ) + cont_t = to_cuda_long("dsv4_c4_prefill_cont", [x for row in cont for x in row]).view(-1, 3) + prev_slot = mapping[self.req_to_token_indexs[cont_t[:, 1], cont_t[:, 2] * 4 - 1].long()].long() + cont_off = (cont_t[:, 2] - 1) % page + assert bool((prev_slot >= 0).all()) and bool(((prev_slot % page) == cont_off).all()) + page_bases[cont_t[:, 0]] = prev_slot - cont_off + + page_idx = page_offsets_t[seg] + torch.div(entries, page, rounding_mode="floor") - first_pages_t[seg] + slots = (page_bases[page_idx] + entries % page).to(torch.int32) + self._register_c4_slots(full_slots, slots) + return + + def _scatter_c4_decode_slots(self, req_list, seq_list, mem_indexes: torch.Tensor) -> None: page = DSV4_C4_PAGE_SIZE mapping = self.mem_manager.full_to_c4_indexs - req_list = b_req_idx_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() mem_indexes = mem_indexes.cuda().long().reshape(-1) cont_rows, cont_prev_pos, cont_offsets = [], [], [] @@ -686,15 +776,11 @@ def _scatter_c4_decode_slots( offsets = torch.tensor(cont_offsets, dtype=torch.int32, device="cuda") assert bool((prev_slots >= 0).all()) assert bool(((prev_slots % page) == (offsets - 1)).all()) - slots = (prev_slots + 1).to(torch.int32) - mapping[mem_indexes[cont_rows]] = slots - self.mem_manager.count_c4_slots(slots, 1) + self._register_c4_slots(mem_indexes[cont_rows], (prev_slots + 1).to(torch.int32)) if new_rows: pages = self.mem_manager.alloc_c4_pages(len(new_rows)).cuda(non_blocking=True).long() - slots = (pages * page).to(torch.int32) - mapping[mem_indexes[new_rows]] = slots - self.mem_manager.count_c4_slots(slots, 1) + self._register_c4_slots(mem_indexes[new_rows], (pages * page).to(torch.int32)) return def _scatter_compress_slots(self, ratio: int, full_slots: torch.Tensor) -> None: @@ -729,11 +815,7 @@ def prepare_prefill_compress_slots( ready_list = b_ready_cache_len_cpu.tolist() seq_list = b_seq_len_cpu.tolist() if self.n_c4 > 0: - for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): - req_idx = int(req_idx) - if req_idx == self.HOLD_REQUEST_ID: - continue - self._scatter_c4_prefill_slots(req_idx, int(ready_len) // 4, int(seq_len) // 4) + self._scatter_c4_prefill_slots_batched(req_list, ready_list, seq_list) if self.n_c128 > 0: ratio = 128 @@ -764,7 +846,7 @@ def prepare_decode_compress_slots( req_list = b_req_idx_cpu.tolist() seq_list = b_seq_len_cpu.tolist() if self.n_c4 > 0: - self._scatter_c4_decode_slots(b_req_idx_cpu, b_seq_len_cpu, mem_indexes) + self._scatter_c4_decode_slots(req_list, seq_list, mem_indexes) if self.n_c128 > 0: ratio = 128 From 82cb6d6a56580dfd83ecda19b1af910427f43ed0 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 24 Jun 2026 11:14:50 +0000 Subject: [PATCH 046/214] fix arguments --- lightllm/models/deepseek_v4/model.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 1be136167a..b3ac16fab3 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -1,5 +1,6 @@ import copy import importlib.util +import json import os import torch @@ -233,6 +234,20 @@ def apply_chat_template( msgs = copy.deepcopy(msgs) + # The model's DSML encoder (encode_arguments_to_dsml in encoding_dsv4.py) expects + # function.arguments as a JSON string and parses it internally. Upstream, + # build_prompt._normalize_tool_call_arguments converts arguments from the OpenAI + # JSON string to a dict (needed by Qwen3.x-style Jinja templates). A dict hits the + # encoder's except-branch and gets wrapped under a single name="arguments" param, + # which the model then imitates and amplifies across turns until required fields go + # missing. Re-serialize dicts back to JSON strings so the encoder emits one + # per real arg. + for msg in msgs: + for tc in msg.get("tool_calls") or []: + fn = tc.get("function") + if isinstance(fn, dict) and isinstance(fn.get("arguments"), dict): + fn["arguments"] = json.dumps(fn["arguments"], ensure_ascii=False) + if tools: wrapped_tools = [] for tool in tools: From bc2259153a9fff7064e4dd9ed2966fb14a0136b1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 24 Jun 2026 11:52:59 +0000 Subject: [PATCH 047/214] tune H100 --- ..._fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json | 110 +++++++++++++ ..._fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json | 110 +++++++++++++ .../{topk_num=6}_NVIDIA_H100_80GB_HBM3.json | 50 ++++++ ...t16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json | 74 +++++++++ ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 74 +++++++++ ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 146 ++++++++++++++++++ 6 files changed, 564 insertions(+) create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..5204097669 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,110 @@ +{ + "12288": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "1536": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "192": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "24576": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "384": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "48": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "49152": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "6": { + "BLOCK_SIZE_K": 64, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "600": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "6144": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "768": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + }, + "96": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 64, + "NEED_TRANS": true, + "num_stages": 2, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..ac4ce1ba57 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,110 @@ +{ + "1": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 4, + "num_warps": 4 + }, + "100": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "1024": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + }, + "128": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 5, + "num_warps": 4 + }, + "16": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 4, + "num_warps": 4 + }, + "2048": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 8 + }, + "256": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 3, + "num_warps": 4 + }, + "32": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 16, + "NEED_TRANS": true, + "num_stages": 5, + "num_warps": 4 + }, + "4096": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": false, + "num_stages": 5, + "num_warps": 8 + }, + "64": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 1, + "NEED_TRANS": true, + "num_stages": 5, + "num_warps": 4 + }, + "8": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "GROUP_SIZE_M": 32, + "NEED_TRANS": true, + "num_stages": 4, + "num_warps": 4 + }, + "8192": { + "BLOCK_SIZE_K": 128, + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "GROUP_SIZE_M": 1, + "NEED_TRANS": false, + "num_stages": 3, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..6aa8d18c54 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,50 @@ +{ + "1": { + "BLOCK_SIZE": 256, + "num_warps": 2 + }, + "100": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "1024": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "128": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "16": { + "BLOCK_SIZE": 512, + "num_warps": 4 + }, + "2048": { + "BLOCK_SIZE": 128, + "num_warps": 4 + }, + "256": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "32": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "4096": { + "BLOCK_SIZE": 256, + "num_warps": 4 + }, + "64": { + "BLOCK_SIZE": 128, + "num_warps": 8 + }, + "8": { + "BLOCK_SIZE": 256, + "num_warps": 2 + }, + "8192": { + "BLOCK_SIZE": 256, + "num_warps": 4 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..e2da8bc968 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_DIM": 256, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + }, + "100": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "1024": { + "BLOCK_DIM": 256, + "BLOCK_M": 1, + "NUM_STAGE": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "16": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "2048": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + }, + "256": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "32": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "4096": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 4, + "num_warps": 2 + }, + "64": { + "BLOCK_DIM": 512, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 4 + }, + "8": { + "BLOCK_DIM": 64, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 2 + }, + "8192": { + "BLOCK_DIM": 1024, + "BLOCK_M": 1, + "NUM_STAGE": 1, + "num_warps": 8 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..588fd4a934 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 8 + }, + "100": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 1, + "num_warps": 4 + }, + "1024": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 3, + "num_warps": 4 + }, + "128": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 1, + "num_warps": 4 + }, + "16": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 4, + "num_warps": 8 + }, + "2048": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 4, + "num_warps": 2 + }, + "256": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 1 + }, + "32": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 2, + "num_warps": 8 + }, + "4096": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 1, + "num_warps": 2 + }, + "64": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 4 + }, + "8": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 8 + }, + "8192": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 3, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..4d7e8f1183 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,146 @@ +{ + "1": { + "BLOCK_M": 128, + "BLOCK_N": 128, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "100": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1024": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "12288": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "1536": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "16": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "192": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "2048": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "24576": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "256": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "32": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "384": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4096": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "48": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "49152": { + "BLOCK_M": 32, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "6": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "600": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "6144": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "768": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "8": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 1, + "num_warps": 8 + }, + "8192": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "96": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + } +} \ No newline at end of file From 77baa054649facd66c8eb09d553692fb2070fc70 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 25 Jun 2026 08:48:30 +0000 Subject: [PATCH 048/214] add encoding_dsv4 --- .../models/deepseek_v4/encoding/__init__.py | 0 .../deepseek_v4/encoding/encoding_dsv4.py | 762 ++++++++++++++++++ lightllm/models/deepseek_v4/model.py | 7 + 3 files changed, 769 insertions(+) create mode 100644 lightllm/models/deepseek_v4/encoding/__init__.py create mode 100644 lightllm/models/deepseek_v4/encoding/encoding_dsv4.py diff --git a/lightllm/models/deepseek_v4/encoding/__init__.py b/lightllm/models/deepseek_v4/encoding/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4/encoding/encoding_dsv4.py b/lightllm/models/deepseek_v4/encoding/encoding_dsv4.py new file mode 100644 index 0000000000..6cbd5f9bfb --- /dev/null +++ b/lightllm/models/deepseek_v4/encoding/encoding_dsv4.py @@ -0,0 +1,762 @@ +""" +DeepSeek-V4 Encoding + +A self-contained implementation for encoding/decoding DeepSeek-V4 chat messages +with tool calling, thinking mode, and quick instruction task support. +""" + +from typing import Any, Dict, List, Union, Optional, Tuple +import copy +import json +import re + +# ============================================================ +# Special Tokens +# ============================================================ + +bos_token: str = "<|begin▁of▁sentence|>" +eos_token: str = "<|end▁of▁sentence|>" +thinking_start_token: str = "" +thinking_end_token: str = "" +dsml_token: str = "|DSML|" + +USER_SP_TOKEN = "<|User|>" +ASSISTANT_SP_TOKEN = "<|Assistant|>" +LATEST_REMINDER_SP_TOKEN = "<|latest_reminder|>" + +# Task special tokens for internal classification tasks +DS_TASK_SP_TOKENS = { + "action": "<|action|>", + "query": "<|query|>", + "authority": "<|authority|>", + "domain": "<|domain|>", + "title": "<|title|>", + "read_url": "<|read_url|>", +} +VALID_TASKS = set(DS_TASK_SP_TOKENS.keys()) + +# ============================================================ +# Templates +# ============================================================ + +system_msg_template: str = "{content}" +user_msg_template: str = "{content}" +latest_reminder_msg_template: str = "{content}" +assistant_msg_template: str = "{reasoning}{content}{tool_calls}" + eos_token +assistant_msg_wo_eos_template: str = "{reasoning}{content}{tool_calls}" +thinking_template: str = "{reasoning_content}" + +response_format_template: str = ( + "## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n{schema}" +) +tool_call_template: str = '<{dsml_token}invoke name="{name}">\n{arguments}\n' +tool_calls_template = "<{dsml_token}{tc_block_name}>\n{tool_calls}\n" +tool_calls_block_name: str = "tool_calls" + +tool_output_template: str = "{content}" + +REASONING_EFFORT_MAX = ( + "Reasoning Effort: Absolute maximum with no shortcuts permitted.\n" + "You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n" # noqa: E501 + "Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n" # noqa: E501 +) + +TOOLS_TEMPLATE = """## Tools + +You have access to a set of tools to help answer the user's question. You can invoke tools by writing a \ +"<{dsml_token}tool_calls>" block like the following: + +<{dsml_token}tool_calls> +<{dsml_token}invoke name="$TOOL_NAME"> +<{dsml_token}parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE +... + +<{dsml_token}invoke name="$TOOL_NAME2"> +... + + + +String parameters should be specified as is and set `string="true"`. For all other types (numbers, \ +booleans, arrays, objects), pass the value in JSON format and set `string="false"`. + +If thinking_mode is enabled (triggered by {thinking_start_token}), you MUST output your complete \ +reasoning inside {thinking_start_token}...{thinking_end_token} BEFORE any tool calls or final response. + +Otherwise, output directly after {thinking_end_token} with tool calls or final response. + +### Available Tool Schemas + +{tool_schemas} + +You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. +""" + +# ============================================================ +# Utility Functions +# ============================================================ + + +def to_json(value: Any) -> str: + """Serialize a value to JSON string.""" + try: + return json.dumps(value, ensure_ascii=False) + except: + return json.dumps(value, ensure_ascii=True) + + +def tools_from_openai_format(tools): + """Extract function definitions from OpenAI-format tool list.""" + return [tool["function"] for tool in tools] + + +def tool_calls_from_openai_format(tool_calls): + """Convert OpenAI-format tool calls to internal format.""" + return [ + { + "name": tool_call["function"]["name"], + "arguments": tool_call["function"]["arguments"], + } + for tool_call in tool_calls + ] + + +def tool_calls_to_openai_format(tool_calls): + """Convert internal tool calls to OpenAI format.""" + return [ + { + "type": "function", + "function": { + "name": tool_call["name"], + "arguments": tool_call["arguments"], + }, + } + for tool_call in tool_calls + ] + + +def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str: + """ + Encode tool call arguments into DSML parameter format. + + Args: + tool_call: Dict with "name" and "arguments" (JSON string) keys. + + Returns: + DSML-formatted parameter string. + """ + p_dsml_template = '<{dsml_token}parameter name="{key}" string="{is_str}">{value}' + P_dsml_strs = [] + + try: + arguments = json.loads(tool_call["arguments"]) + except Exception: + arguments = {"arguments": tool_call["arguments"]} + + for k, v in arguments.items(): + p_dsml_str = p_dsml_template.format( + dsml_token=dsml_token, + key=k, + is_str="true" if isinstance(v, str) else "false", + value=v if isinstance(v, str) else to_json(v), + ) + P_dsml_strs.append(p_dsml_str) + + return "\n".join(P_dsml_strs) + + +def decode_dsml_to_arguments(tool_name: str, tool_args: Dict[str, Tuple[str, str]]) -> Dict[str, str]: + """ + Decode DSML parameters back to a tool call dict. + + Args: + tool_name: Name of the tool. + tool_args: Dict mapping param_name -> (value, is_string_flag). + + Returns: + Dict with "name" and "arguments" (JSON string) keys. + """ + + def _decode_value(key: str, value: str, string: str): + if string == "true": + value = to_json(value) + return f"{to_json(key)}: {value}" + + tool_args_json = "{" + ", ".join([_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()]) + "}" + return dict(name=tool_name, arguments=tool_args_json) + + +def render_tools(tools: List[Dict[str, Union[str, Dict[str, Any]]]]) -> str: + """ + Render tool schemas into the system prompt format. + + Args: + tools: List of tool schema dicts (each with name, description, parameters). + + Returns: + Formatted tools section string. + """ + tools_json = [to_json(t) for t in tools] + + return TOOLS_TEMPLATE.format( + tool_schemas="\n".join(tools_json), + dsml_token=dsml_token, + thinking_start_token=thinking_start_token, + thinking_end_token=thinking_end_token, + ) + + +def find_last_user_index(messages: List[Dict[str, Any]]) -> int: + """Find the index of the last user/developer message.""" + last_user_index = -1 + for idx in range(len(messages) - 1, -1, -1): + if messages[idx].get("role") in ["user", "developer"]: + last_user_index = idx + break + return last_user_index + + +# ============================================================ +# Message Rendering +# ============================================================ + + +def render_message( + index: int, + messages: List[Dict[str, Any]], + thinking_mode: str, + drop_thinking: bool = True, + reasoning_effort: Optional[str] = None, +) -> str: + """ + Render a single message at the given index into its encoded string form. + + This is the core function that converts each message in the conversation + into the DeepSeek-V4 format. + + Args: + index: Index of the message to render. + messages: Full list of messages in the conversation. + thinking_mode: Either "chat" or "thinking". + drop_thinking: Whether to drop reasoning content from earlier turns. + reasoning_effort: Optional reasoning effort level ("max", "high", or None). + + Returns: + Encoded string for this message. + """ + assert 0 <= index < len(messages) + assert thinking_mode in ["chat", "thinking"], f"Invalid thinking_mode `{thinking_mode}`" + + prompt = "" + msg = messages[index] + last_user_idx = find_last_user_index(messages) + + role = msg.get("role") + content = msg.get("content") + tools = msg.get("tools") + response_format = msg.get("response_format") + tool_calls = msg.get("tool_calls") + reasoning_content = msg.get("reasoning_content") + wo_eos = msg.get("wo_eos", False) + + if tools: + tools = tools_from_openai_format(tools) + if tool_calls: + tool_calls = tool_calls_from_openai_format(tool_calls) + + # Reasoning effort prefix (only at index 0 in thinking mode with max effort) + assert reasoning_effort in ["max", None, "high"], f"Invalid reasoning effort: {reasoning_effort}" + if index == 0 and thinking_mode == "thinking" and reasoning_effort == "max": + prompt += REASONING_EFFORT_MAX + + if role == "system": + prompt += system_msg_template.format(content=content or "") + if tools: + prompt += "\n\n" + render_tools(tools) + if response_format: + prompt += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + elif role == "developer": + assert content, f"Invalid message for role `{role}`: {msg}" + + content_developer = USER_SP_TOKEN + content_developer += content + + if tools: + content_developer += "\n\n" + render_tools(tools) + if response_format: + content_developer += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + prompt += user_msg_template.format(content=content_developer) + + elif role == "user": + prompt += USER_SP_TOKEN + + # Handle content blocks (tool results mixed with text) + content_blocks = msg.get("content_blocks") + if content_blocks: + parts = [] + for block in content_blocks: + block_type = block.get("type") + if block_type == "text": + parts.append(block.get("text", "")) + elif block_type == "tool_result": + tool_content = block.get("content", "") + if isinstance(tool_content, list): + text_parts = [] + for b in tool_content: + if b.get("type") == "text": + text_parts.append(b.get("text", "")) + else: + text_parts.append(f"[Unsupported {b.get('type')}]") + tool_content = "\n\n".join(text_parts) + parts.append(tool_output_template.format(content=tool_content)) + else: + parts.append(f"[Unsupported {block_type}]") + prompt += "\n\n".join(parts) + else: + prompt += content or "" + + elif role == "latest_reminder": + prompt += LATEST_REMINDER_SP_TOKEN + latest_reminder_msg_template.format(content=content) + + elif role == "tool": + raise NotImplementedError( + "deepseek_v4 merges tool messages into user; please preprocess with merge_tool_messages()" + ) + + elif role == "assistant": + thinking_part = "" + tc_content = "" + + if tool_calls: + tc_list = [ + tool_call_template.format( + dsml_token=dsml_token, name=tc.get("name"), arguments=encode_arguments_to_dsml(tc) + ) + for tc in tool_calls + ] + tc_content += "\n\n" + tool_calls_template.format( + dsml_token=dsml_token, + tool_calls="\n".join(tc_list), + tc_block_name=tool_calls_block_name, + ) + + summary_content = content or "" + rc = reasoning_content or "" + + # Check if previous message has a task - if so, this is a task output (no thinking) + prev_has_task = index - 1 >= 0 and messages[index - 1].get("task") is not None + + if thinking_mode == "thinking" and not prev_has_task: + if not drop_thinking or index > last_user_idx: + thinking_part = thinking_template.format(reasoning_content=rc) + thinking_end_token + else: + thinking_part = "" + + if wo_eos: + prompt += assistant_msg_wo_eos_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + prompt += assistant_msg_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + raise NotImplementedError(f"Unknown role: {role}") + + # Append transition tokens based on what follows + if index + 1 < len(messages) and messages[index + 1].get("role") not in ["assistant", "latest_reminder"]: + return prompt + + task = messages[index].get("task") + if task is not None: + # Task special token for internal classification tasks + assert task in VALID_TASKS, f"Invalid task: '{task}'. Valid tasks are: {list(VALID_TASKS)}" + task_sp_token = DS_TASK_SP_TOKENS[task] + + if task != "action": + # Non-action tasks: append task sp token directly after the message + prompt += task_sp_token + else: + # Action task: append Assistant + thinking token + action sp token + prompt += ASSISTANT_SP_TOKEN + prompt += thinking_end_token if thinking_mode != "thinking" else thinking_start_token + prompt += task_sp_token + + elif messages[index].get("role") in ["user", "developer"]: + # Normal generation: append Assistant + thinking token + prompt += ASSISTANT_SP_TOKEN + if not drop_thinking and thinking_mode == "thinking": + prompt += thinking_start_token + elif drop_thinking and thinking_mode == "thinking" and index >= last_user_idx: + prompt += thinking_start_token + else: + prompt += thinking_end_token + + return prompt + + +# ============================================================ +# Preprocessing +# ============================================================ + + +def merge_tool_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Merge tool messages into the preceding user message using content_blocks format. + + DeepSeek-V4 does not have a standalone "tool" role; instead, tool results + are encoded as blocks within user messages. + + This function converts a standard OpenAI-format conversation (with separate + "tool" role messages) into V4 format where tool results are merged into + user messages. + + Args: + messages: List of message dicts in OpenAI format. + + Returns: + Processed message list with tool messages merged into user messages. + """ + merged: List[Dict[str, Any]] = [] + + for msg in messages: + msg = copy.deepcopy(msg) + role = msg.get("role") + + if role == "tool": + # Convert tool message to a user message with tool_result block + tool_block = { + "type": "tool_result", + "tool_use_id": msg.get("tool_call_id", ""), + "content": msg.get("content", ""), + } + # Merge into previous message if it's already a user (merged tool) + if merged and merged[-1].get("role") == "user" and "content_blocks" in merged[-1]: + merged[-1]["content_blocks"].append(tool_block) + else: + merged.append( + { + "role": "user", + "content_blocks": [tool_block], + } + ) + elif role == "user": + text_block = {"type": "text", "text": msg.get("content", "")} + if ( + merged + and merged[-1].get("role") == "user" + and "content_blocks" in merged[-1] + and merged[-1].get("task") is None + ): + merged[-1]["content_blocks"].append(text_block) + else: + new_msg = { + "role": "user", + "content": msg.get("content", ""), + "content_blocks": [text_block], + } + # Preserve extra fields (task, wo_eos, mask, etc.) + for key in ("task", "wo_eos", "mask"): + if key in msg: + new_msg[key] = msg[key] + merged.append(new_msg) + else: + merged.append(msg) + + return merged + + +def sort_tool_results_by_call_order(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Sort tool_result blocks within user messages by the order of tool_calls + in the preceding assistant message. + + Args: + messages: Preprocessed message list (after merge_tool_messages). + + Returns: + Message list with sorted tool result blocks. + """ + last_tool_call_order: Dict[str, int] = {} + + for msg in messages: + role = msg.get("role") + if role == "assistant" and msg.get("tool_calls"): + last_tool_call_order = {} + for idx, tc in enumerate(msg["tool_calls"]): + tc_id = tc.get("id") or tc.get("function", {}).get("id", "") + if tc_id: + last_tool_call_order[tc_id] = idx + + elif role == "user" and msg.get("content_blocks"): + tool_blocks = [b for b in msg["content_blocks"] if b.get("type") == "tool_result"] + if len(tool_blocks) > 1 and last_tool_call_order: + sorted_blocks = sorted(tool_blocks, key=lambda b: last_tool_call_order.get(b.get("tool_use_id", ""), 0)) + sorted_idx = 0 + new_blocks = [] + for block in msg["content_blocks"]: + if block.get("type") == "tool_result": + new_blocks.append(sorted_blocks[sorted_idx]) + sorted_idx += 1 + else: + new_blocks.append(block) + msg["content_blocks"] = new_blocks + + return messages + + +# ============================================================ +# Main Encoding Function +# ============================================================ + + +def encode_messages( + messages: List[Dict[str, Any]], + thinking_mode: str, + context: Optional[List[Dict[str, Any]]] = None, + drop_thinking: bool = True, + add_default_bos_token: bool = True, + reasoning_effort: Optional[str] = None, +) -> str: + """ + Encode a list of messages into the DeepSeek-V4 prompt format. + + This is the main entry point for encoding conversations. It handles: + - BOS token insertion + - Thinking mode with optional reasoning content dropping + - Tool message merging into user messages + - Multi-turn conversation context + + Args: + messages: List of message dicts to encode. + thinking_mode: Either "chat" or "thinking". + context: Optional preceding context messages (already encoded prefix). + drop_thinking: If True, drop reasoning_content from earlier assistant turns + (only keep reasoning for messages after the last user message). + add_default_bos_token: Whether to prepend BOS token at conversation start. + reasoning_effort: Optional reasoning effort level ("max", "high", or None). + + Returns: + The encoded prompt string. + """ + context = context if context else [] + + # Preprocess: merge tool messages and sort tool results + messages = merge_tool_messages(messages) + messages = sort_tool_results_by_call_order(context + messages)[len(context) :] + if context: + context = merge_tool_messages(context) + context = sort_tool_results_by_call_order(context) + + full_messages = context + messages + + prompt = bos_token if add_default_bos_token and len(context) == 0 else "" + + # Resolve drop_thinking: if any message has tools defined, don't drop thinking + effective_drop_thinking = drop_thinking + if any(m.get("tools") for m in full_messages): + effective_drop_thinking = False + + if thinking_mode == "thinking" and effective_drop_thinking: + full_messages = _drop_thinking_messages(full_messages) + # After dropping, recalculate how many messages to render + # (context may have shrunk too) + num_to_render = len(full_messages) - len(_drop_thinking_messages(context)) + context_len = len(full_messages) - num_to_render + else: + num_to_render = len(messages) + context_len = len(context) + + for idx in range(num_to_render): + prompt += render_message( + idx + context_len, + full_messages, + thinking_mode=thinking_mode, + drop_thinking=effective_drop_thinking, + reasoning_effort=reasoning_effort, + ) + + return prompt + + +def _drop_thinking_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Drop reasoning_content and non-essential messages before the last user message. + + Behavior: + - Messages with role in ["user", "system", "tool", "latest_reminder"] are always kept. + - Messages at or after the last user index are always kept. + - Assistant messages before the last user get reasoning_content removed. + - Developer messages before the last user are dropped entirely. + """ + last_user_idx = find_last_user_index(messages) + result = [] + keep_roles = {"user", "system", "tool", "latest_reminder", "direct_search_results"} + + for idx, msg in enumerate(messages): + role = msg.get("role") + if role in keep_roles or idx >= last_user_idx: + result.append(msg) + elif role == "assistant": + msg = copy.copy(msg) + msg.pop("reasoning_content", None) + result.append(msg) + # developer and other roles before last_user_idx are dropped + + return result + + +# ============================================================ +# Parsing (Decoding model output) +# ============================================================ + + +def _read_until_stop(index: int, text: str, stop: List[str]) -> Tuple[int, str, Optional[str]]: + """ + Read text from index until one of the stop strings is found. + + Returns: + Tuple of (new_index, content_before_stop, matched_stop_string_or_None). + """ + min_pos = len(text) + matched_stop = None + + for s in stop: + pos = text.find(s, index) + if pos != -1 and pos < min_pos: + min_pos = pos + matched_stop = s + + if matched_stop: + content = text[index:min_pos] + return min_pos + len(matched_stop), content, matched_stop + else: + content = text[index:] + return len(text), content, None + + +def parse_tool_calls(index: int, text: str) -> Tuple[int, Optional[str], List[Dict[str, str]]]: + """ + Parse DSML tool calls from text starting at the given index. + + Args: + index: Starting position in text. + text: The full text to parse. + + Returns: + Tuple of (new_index, last_stop_token, list_of_tool_call_dicts). + Each tool call dict has "name" and "arguments" keys. + """ + tool_calls: List[Dict[str, Any]] = [] + stop_token = None + tool_calls_end_token = f"" + + while index < len(text): + index, _, stop_token = _read_until_stop(index, text, [f"<{dsml_token}invoke", tool_calls_end_token]) + if _ != ">\n": + raise ValueError(f"Tool call format error: expected '>\\n' but got '{_}'") + + if stop_token == tool_calls_end_token: + break + + if stop_token is None: + raise ValueError("Missing special token in tool calls") + + index, tool_name_content, stop_token = _read_until_stop( + index, text, [f"<{dsml_token}parameter", f"\n$', tool_name_content, flags=re.DOTALL) + if len(p_tool_name) != 1: + raise ValueError(f"Tool name format error: '{tool_name_content}'") + tool_name = p_tool_name[0] + + tool_args: Dict[str, Tuple[str, str]] = {} + while stop_token == f"<{dsml_token}parameter": + index, param_content, stop_token = _read_until_stop(index, text, [f"/{dsml_token}parameter"]) + + param_kv = re.findall(r'^ name="(.*?)" string="(true|false)">(.*?)<$', param_content, flags=re.DOTALL) + if len(param_kv) != 1: + raise ValueError(f"Parameter format error: '{param_content}'") + param_name, string, param_value = param_kv[0] + + if param_name in tool_args: + raise ValueError(f"Duplicate parameter name: '{param_name}'") + tool_args[param_name] = (param_value, string) + + index, content, stop_token = _read_until_stop( + index, text, [f"<{dsml_token}parameter", f"\n": + raise ValueError(f"Parameter format error: expected '>\\n' but got '{content}'") + + tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) + tool_calls.append(tool_call) + + return index, stop_token, tool_calls + + +def parse_message_from_completion_text(text: str, thinking_mode: str) -> Dict[str, Any]: + """ + Parse a model completion text into a structured assistant message. + + This function takes the raw text output from the model (a single assistant turn) + and extracts: + - reasoning_content (thinking block) + - content (summary/response) + - tool_calls (if any) + + NOTE: This function is designed to parse only correctly formatted strings and + will raise ValueError for malformed output. + + Args: + text: The raw completion text (including EOS token). + thinking_mode: Either "chat" or "thinking". + + Returns: + Dict with keys: "role", "content", "reasoning_content", "tool_calls". + tool_calls are in OpenAI format. + """ + summary_content, reasoning_content, tool_calls = "", "", [] + index, stop_token = 0, None + tool_calls_start_token = f"\n\n<{dsml_token}{tool_calls_block_name}" + + is_thinking = thinking_mode == "thinking" + is_tool_calling = False + + if is_thinking: + index, content_delta, stop_token = _read_until_stop(index, text, [thinking_end_token, tool_calls_start_token]) + reasoning_content = content_delta + assert stop_token == thinking_end_token, "Invalid thinking format: missing " + + index, content_delta, stop_token = _read_until_stop(index, text, [eos_token, tool_calls_start_token]) + summary_content = content_delta + if stop_token == tool_calls_start_token: + is_tool_calling = True + else: + assert stop_token == eos_token, "Invalid format: missing EOS token" + + if is_tool_calling: + index, stop_token, tool_calls = parse_tool_calls(index, text) + + index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) + assert not tool_ends_text, "Unexpected content after tool calls" + + assert len(text) == index and stop_token in [eos_token, None], "Unexpected content at end" + + for sp_token in [bos_token, eos_token, thinking_start_token, thinking_end_token, dsml_token]: + assert ( + sp_token not in summary_content and sp_token not in reasoning_content + ), f"Unexpected special token '{sp_token}' in content" + + return { + "role": "assistant", + "content": summary_content, + "reasoning_content": reasoning_content, + "tool_calls": tool_calls_to_openai_format(tool_calls), + } diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index b3ac16fab3..5353dfcade 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -205,7 +205,14 @@ def _get_encoding_module(self): if self._encoding_module is not None: return self._encoding_module + # Prefer the encoder shipped inside the model dir (respects any model-specific + # customization); fall back to the copy vendored in this repo, because some + # DeepSeek-V4 releases (e.g. the FP8 weights) do NOT ship an encoding/ dir. + # vLLM/sglang likewise vendor this encoder in-tree instead of depending on the + # model directory. encoding_path = os.path.join(self.model_dir, "encoding", "encoding_dsv4.py") + if not os.path.exists(encoding_path): + encoding_path = os.path.join(os.path.dirname(__file__), "encoding", "encoding_dsv4.py") if not os.path.exists(encoding_path): raise FileNotFoundError(f"DeepSeek-V4 encoding file not found: {encoding_path}") From 2225e1aa115601807c93821aea4f7bb17bd96003 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 26 Jun 2026 05:36:27 +0000 Subject: [PATCH 049/214] fix c4 error --- lightllm/common/req_manager.py | 30 ++++++++++ .../router/dynamic_prompt/radix_cache.py | 41 +++++++++++++ .../server/router/model_infer/infer_batch.py | 60 +++++++++++++++++++ .../model_infer/mode_backend/base_backend.py | 32 ++++++++-- 4 files changed, 159 insertions(+), 4 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index d40fe96989..1af4991a3d 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -675,6 +675,7 @@ def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list) -> N # host plan (cheap int arithmetic, no GPU). dup req in one call would break vectorized # continuation ordering -> fall back to the safe per-req path. plan, seen, duplicate_req = [], set(), False + c4_page_need = 0 for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): req_idx = int(req_idx) if req_idx == self.HOLD_REQUEST_ID: @@ -682,11 +683,14 @@ def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list) -> N first, last = int(ready_len) // 4, int(seq_len) // 4 if last <= first: continue + c4_page_need += (last - 1) // page - first // page + 1 # 上界=区间触及页数, 复用本循环 duplicate_req |= req_idx in seen seen.add(req_idx) plan.append((req_idx, first, last)) if not plan: return + # 兑现: 在所有分支(dup/fresh/batched)的 alloc_c4_pages 之前统一腾页 + self._realize_c4_pages(c4_page_need) if duplicate_req: for req_idx, first, last in plan: self._scatter_c4_prefill_slots(req_idx, first, last) @@ -779,6 +783,7 @@ def _scatter_c4_decode_slots(self, req_list, seq_list, mem_indexes: torch.Tensor self._register_c4_slots(mem_indexes[cont_rows], (prev_slots + 1).to(torch.int32)) if new_rows: + self._realize_c4_pages(len(new_rows)) # 兑现: 精确需求, 复用已算的 new_rows pages = self.mem_manager.alloc_c4_pages(len(new_rows)).cuda(non_blocking=True).long() self._register_c4_slots(mem_indexes[new_rows], (pages * page).to(torch.int32)) return @@ -793,10 +798,35 @@ def _scatter_compress_slots(self, ratio: int, full_slots: torch.Tensor) -> None: need = torch.unique(full_slots[mapping[full_slots] < 0]) if need.numel() == 0: return + if ratio == 128: # _scatter_compress_slots 仅用于 c128; 兑现其槽(复用已算的 need.numel()) + self._realize_c128_slots(int(need.numel())) new_slots = alloc(need.numel()).cuda(non_blocking=True).to(torch.int32) mapping[need] = new_slots return + def _realize_c4_pages(self, need_pages: int) -> None: + """压缩池兑现 —— 和主池在 prep 里调 free_radix_cache_to_get_enough_token 同一套路: + base_backend admission 已按"空闲+可回收"放行本步请求,这里在真分配前(scatter 已算好 need) + 把可回收的无引用 radix 节点驱逐出来腾出 c4 页,避免 alloc_c4_pages 触底 assert。 + 可回收仍不足时由 admission 的 wait_pause 兜底。""" + if self.n_c4 == 0 or need_pages <= 0: + return + # 延迟 import: infer_batch 在模块顶 import 了 req_manager,顶层 import 会循环引用 + 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_c4_pages(need_pages) + return + + def _realize_c128_slots(self, need_slots: int) -> None: + if self.n_c128 == 0 or need_slots <= 0: + return + 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_c128_slots(need_slots) + return + def prepare_prefill_compress_slots( self, b_req_idx: torch.Tensor, diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index dbffea0a33..e45b556b57 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -5,6 +5,9 @@ from typing import Any, Tuple, Dict, Set, List, Optional, Union from sortedcontainers import SortedSet from .shared_arr import SharedArray +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) class UniqueTimeIdGenerator: @@ -703,6 +706,44 @@ def release_mem(mem_index): self.mem_manager.free(mem_index) return + def _free_radix_full_nodes_until(self, allocator, need: int) -> None: + """DeepSeek-V4 压缩池(c4/c128)兑现: 沿 LRU 序逐个驱逐 ref_count==0 的整个 full radix 节点, + 经 mem_manager.free() 级联回收其 c4 页 / c128 槽(evict_c4/evict_c128),每驱逐一个就复查 + *真实* allocator(不靠计数,稳),直到够或已无可驱逐的无引用节点。后者(空闲+可回收仍不足) + 由上游 base_backend admission 的 wait_pause 兜底,allocator 的 assert 是最后防线。""" + if self.mem_manager is None or allocator is None: + return + while allocator.can_use_mem_size < need: + # 无可驱逐的无引用 token => 停(admission 应已 wait_pause) + if self.tree_total_tokens_num.arr[0] <= self.refed_tokens_num.arr[0]: + # 兜底没兜住:admission/realize 估算漂移了。打日志便于定位(否则只会撞下游隐晦的 + # allocator "error alloc state" assert)。 + logger.warning( + f"dsv4 compress-pool realize could not free enough: need={need} " + f"free={allocator.can_use_mem_size} tree_total={self.tree_total_tokens_num.arr[0]} " + f"refed={self.refed_tokens_num.arr[0]} (admission should have paused this req)" + ) + return + release_mems = [] + # 复用已测的 evict():弹一个 LRU、ref==0 的叶子(>=1 token),其 full 槽经 free 级联回收压缩槽 + self.evict(1, lambda mem_index: release_mems.append(mem_index)) + self.mem_manager.free(torch.concat(release_mems)) + return + + def free_radix_cache_to_get_enough_c4_pages(self, need_pages: int) -> None: + allocator = getattr(self.mem_manager, "c4_page_allocator", None) if self.mem_manager is not None else None + if allocator is None or need_pages <= 0: + return + self._free_radix_full_nodes_until(allocator, need_pages) + return + + def free_radix_cache_to_get_enough_c128_slots(self, need_slots: int) -> None: + allocator = getattr(self.mem_manager, "c128_allocator", None) if self.mem_manager is not None else None + if allocator is None or need_slots <= 0: + return + self._free_radix_full_nodes_until(allocator, need_slots) + return + class _RadixCacheReadOnlyClient: """ diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index f73fe6cad4..1d5db33bfb 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -481,6 +481,66 @@ def _get_dsv4_swa_page_size(self): return None return mem_manager.swa_pool.page_size + # ---- DeepSeek-V4 compressed-pool (c4/c128) admission, mirror of the swa helpers above ---- + # c4 is paged for fp8_paged_mqa_logits: 64 c4-slots/page == 256 full tokens. The prompt-cache + # radix is 256-aligned (DSV4_PROMPT_CACHE_PAGE_SIZE) and c4 is NOT windowed, so reclaimable c4 + # pages derive exactly from the unref token count (`// 256`) — no separate counter needed + # (unlike swa, which is windowed). c128 is slot-based: 1 slot per 128 full tokens (`// 128`). + def get_can_alloc_dsv4_c4_page_num(self): + allocator = getattr(self.req_manager.mem_manager, "c4_page_allocator", None) + if allocator is None: + return None + radix_unref_page_num = 0 + if self.radix_cache is not None: + radix_unref_page_num = ( + self.radix_cache.get_tree_total_tokens_num() - self.radix_cache.get_refed_tokens_num() + ) // 256 + return int(allocator.can_use_mem_size) + int(radix_unref_page_num) + + def get_can_alloc_dsv4_c128_slot_num(self): + allocator = getattr(self.req_manager.mem_manager, "c128_allocator", None) + if allocator is None: + return None + radix_unref_slot_num = 0 + if self.radix_cache is not None: + radix_unref_slot_num = ( + self.radix_cache.get_tree_total_tokens_num() - self.radix_cache.get_refed_tokens_num() + ) // 128 + return int(allocator.can_use_mem_size) + int(radix_unref_slot_num) + + def get_dsv4_c4_decode_need_page_num(self, req: "InferReq"): + if getattr(self.req_manager.mem_manager, "c4_page_allocator", None) is None: + return 0 + seq_len = int(req.get_cur_total_len()) + # 与 _scatter_c4_decode_slots 一致: 关组(seq%4==0)且组末 c4-entry 落页首(entry%64==0) -> 开新页 + if seq_len > 0 and seq_len % 4 == 0 and (seq_len // 4 - 1) % 64 == 0: + return 1 + return 0 + + def get_dsv4_c128_decode_need_slot_num(self, req: "InferReq"): + if getattr(self.req_manager.mem_manager, "c128_allocator", None) is None: + return 0 + seq_len = int(req.get_cur_total_len()) + return 1 if (seq_len > 0 and seq_len % 128 == 0) else 0 + + def get_dsv4_c4_prefill_need_page_num(self, req: "InferReq", is_chuncked_prefill: bool): + if getattr(self.req_manager.mem_manager, "c4_page_allocator", None) is None: + return 0 + start = int(req.cur_kv_len) + end = int(req.get_chuncked_input_token_len()) if is_chuncked_prefill else int(req.get_cur_total_len()) + first, last = start // 4, end // 4 + if last <= first: + return 0 + # 安全上界: 覆盖 c4-entry 区间 [first,last) 触及的全部 64-页(忽略已分配的延续页 -> 偏多, 安全) + return (last - 1) // 64 - first // 64 + 1 + + def get_dsv4_c128_prefill_need_slot_num(self, req: "InferReq", is_chuncked_prefill: bool): + if getattr(self.req_manager.mem_manager, "c128_allocator", None) is None: + return 0 + start = int(req.cur_kv_len) + end = int(req.get_chuncked_input_token_len()) if is_chuncked_prefill else int(req.get_cur_total_len()) + return max(0, end // 128 - start // 128) + def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """ 该函数用于在线性混合模型prefill后,如果存在大页匹配的情况下,将线性层状态复制到 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 3319a42db5..07eaeefa14 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -598,6 +598,8 @@ def _get_classed_reqs( can_alloc_token_num = g_infer_context.get_can_alloc_token_num() can_alloc_dsv4_swa_page_num = g_infer_context.get_can_alloc_dsv4_swa_page_num() + can_alloc_dsv4_c4_page_num = g_infer_context.get_can_alloc_dsv4_c4_page_num() + can_alloc_dsv4_c128_slot_num = g_infer_context.get_can_alloc_dsv4_c128_slot_num() for req_obj in ready_reqs: @@ -632,13 +634,22 @@ def _get_classed_reqs( if is_decode: token_num = req_obj.decode_need_token_num() swa_page_num = g_infer_context.get_dsv4_swa_decode_need_page_num(req_obj) - if token_num <= can_alloc_token_num and ( - can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num + c4_page_num = g_infer_context.get_dsv4_c4_decode_need_page_num(req_obj) + c128_slot_num = g_infer_context.get_dsv4_c128_decode_need_slot_num(req_obj) + if ( + token_num <= can_alloc_token_num + and (can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num) + and (can_alloc_dsv4_c4_page_num is None or c4_page_num <= can_alloc_dsv4_c4_page_num) + and (can_alloc_dsv4_c128_slot_num is None or c128_slot_num <= can_alloc_dsv4_c128_slot_num) ): decode_reqs.append(req_obj) can_alloc_token_num -= token_num if can_alloc_dsv4_swa_page_num is not None: can_alloc_dsv4_swa_page_num -= swa_page_num + if can_alloc_dsv4_c4_page_num is not None: + can_alloc_dsv4_c4_page_num -= c4_page_num + if can_alloc_dsv4_c128_slot_num is not None: + can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True @@ -656,14 +667,27 @@ def _get_classed_reqs( swa_page_num = g_infer_context.get_dsv4_swa_prefill_need_page_num( req_obj, is_chuncked_prefill=not self.disable_chunked_prefill ) - if token_num <= can_alloc_token_num and ( - can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num + c4_page_num = g_infer_context.get_dsv4_c4_prefill_need_page_num( + req_obj, is_chuncked_prefill=not self.disable_chunked_prefill + ) + c128_slot_num = g_infer_context.get_dsv4_c128_prefill_need_slot_num( + req_obj, is_chuncked_prefill=not self.disable_chunked_prefill + ) + if ( + token_num <= can_alloc_token_num + and (can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num) + and (can_alloc_dsv4_c4_page_num is None or c4_page_num <= can_alloc_dsv4_c4_page_num) + and (can_alloc_dsv4_c128_slot_num is None or c128_slot_num <= can_alloc_dsv4_c128_slot_num) ): prefill_tokens += token_num prefill_reqs.append(req_obj) can_alloc_token_num -= token_num if can_alloc_dsv4_swa_page_num is not None: can_alloc_dsv4_swa_page_num -= swa_page_num + if can_alloc_dsv4_c4_page_num is not None: + can_alloc_dsv4_c4_page_num -= c4_page_num + if can_alloc_dsv4_c128_slot_num is not None: + can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True From 112247b2ff3d779993db83b543cad7298aee87f6 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 26 Jun 2026 05:44:30 +0000 Subject: [PATCH 050/214] fuse wq_a+wkv & indexer wkv+wgate GEMMs; fp8 wo_a at tp8 (1 group/rank) --- .../layer_infer/transformer_layer_infer.py | 25 ++++-- .../layer_weights/transformer_layer_weight.py | 86 ++++++++++--------- 2 files changed, 61 insertions(+), 50 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 464639f0f7..d8b4955d6f 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -150,7 +150,10 @@ def _get_qkv( input = self._tpsp_allgather(input=input, infer_state=infer_state) T = input.shape[0] - qa = layer_weight.q_norm_(layer_weight.wq_a_.mm(input), eps=self.eps_) + # wq_a and wkv share `input` -> one fused fp8 GEMM, split [q_lora_rank | head_dim]. qa is a + # row-strided view (rmsnorm honors stride(0)); kv feeds a sglang jit kernel -> contiguous. + qkv = layer_weight.wq_a_wkv_.mm(input) + qa = layer_weight.q_norm_(qkv[:, : -self.head_dim_], eps=self.eps_) q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, # fused in one sglang dsv4 jit kernel (fp32 norm/rotation, bf16 in between -- same as eager). @@ -162,7 +165,7 @@ def _get_qkv( infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, mem_index=infer_state.mem_index, - kv=layer_weight.wkv_.mm(input), + kv=qkv[:, -self.head_dim_ :].contiguous(), kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, freqs_cis=self.freqs_cis, @@ -175,8 +178,12 @@ def _get_o(self, o, infer_state: DeepseekV4InferStateInfo, layer_weight: Deepsee position_cos, position_sin = self._select_rope(infer_state) rotary_emb_fwd(o[..., -self.qk_rope_head_dim :], None, position_cos, position_sin, inverse=True) T = o.shape[0] - o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] - o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] + if layer_weight.o_proj_fp8: + # one group per rank -> a single fp8 GEMM (deepgemm .mm quantizes o to fp8 internally) + o = layer_weight.wo_a_.mm(o.reshape(T, -1)) # [T, o_lora] + else: + o = o.reshape(T, self.tp_groups, -1).transpose(0, 1).contiguous() # [groups, T, per_group_in] + o = layer_weight.wo_a_.bmm(o).transpose(0, 1).reshape(T, -1) # [T, groups*o_lora] o = layer_weight.wo_b_.mm(o) return self._tpsp_reduce(input=o, infer_state=infer_state) @@ -429,11 +436,11 @@ def prepare_states( ) if self._metadata is not None: if self.is_in_indexer: - # indexer wkv/wgate are two separate replicated weights; cat -> [T, 2*coff*idx_hd] - # (same [kv | score] layout the fused compressor_wkv_gate_ produces for attention). - kv = layer_weight.idx_cmp_wkv_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager) - gate = layer_weight.idx_cmp_wgate_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager) - self._metadata.kv_score = torch.cat([kv, gate], dim=-1).float() + # fused wkv/wgate GEMM -> [T, 2*coff*idx_hd] in the [kv | score] layout directly + # (same as the attention compressor_wkv_gate_). + self._metadata.kv_score = layer_weight.idx_cmp_wkv_gate_.mm( + x, use_custom_tensor_mananger=use_custom_tensor_manager + ).float() ape = layer_weight.idx_cmp_ape_.weight else: self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm(x).float() diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index e818351c2e..5ecce0a763 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -62,11 +62,13 @@ def _init_weight(self): # ------------------------------------------------------------------ attention def _init_qkvo(self): p = f"{self.prefix}.attn" - # q low-rank (a replicated, b column-parallel over heads), kv single head (replicated) - self.wq_a_ = ROWMMWeight( + # q low-rank A and kv (single replicated head) both consume the same attention input -> + # fuse into one fp8 GEMM; _get_qkv splits the [q_lora_rank | head_dim] output. (q_b is + # column-parallel over heads.) + self.wq_a_wkv_ = ROWMMWeight( in_dim=self.hidden, - out_dims=[self.q_lora_rank], - weight_names=f"{p}.wq_a.weight", + out_dims=[self.q_lora_rank, self.head_dim], + weight_names=[f"{p}.wq_a.weight", f"{p}.wkv.weight"], data_type=self.data_type_, quant_method=self.get_quant_method("wq_a"), tp_rank=0, @@ -79,31 +81,37 @@ def _init_qkvo(self): data_type=self.data_type_, quant_method=self.get_quant_method("wq_b"), ) - self.wkv_ = ROWMMWeight( - in_dim=self.hidden, - out_dims=[self.head_dim], - weight_names=f"{p}.wkv.weight", - data_type=self.data_type_, - quant_method=self.get_quant_method("wkv"), - tp_rank=0, - tp_world_size=1, - ) self.q_norm_ = RMSNormWeight(dim=self.q_lora_rank, weight_name=f"{p}.q_norm.weight", data_type=self.data_type_) self.kv_norm_ = RMSNormWeight(dim=self.head_dim, weight_name=f"{p}.kv_norm.weight", data_type=self.data_type_) self.attn_sink_ = TpAttSinkWeight( all_q_head_num=self.n_heads, weight_name=f"{p}.attn_sink", data_type=torch.float32 ) - # grouped low-rank output projection: wo_a is a per-group batched matmul [groups, in, o_lora], - # wo_b is row-parallel [groups*o_lora -> hidden]. wo_a is reshaped in load_hf_weights. + # grouped low-rank output projection (wo_a per-group [in, o_lora], wo_b row-parallel + # [groups*o_lora -> hidden]). per_group_in = self.n_heads * self.head_dim // self.o_groups - self.wo_a_ = ROWBMMWeight( - dim0=self.o_groups, - dim1=per_group_in, - dim2=self.o_lora_rank, - weight_names=f"{p}.wo_a.weight", - data_type=self.data_type_, - quant_method=None, - ) + # When o_groups == tp_world_size (e.g. the daily tp8 config) each rank owns exactly ONE + # group, so the grouped O-proj collapses to a single GEMM -> run it in fp8 (deepgemm) + # instead of dequantizing wo_a to bf16. sglang does the same (fp8 wo_a is default-on there). + # For >1 group per rank (tp < o_groups) the per-group inputs differ (block-diagonal), so + # keep the bf16 grouped bmm. + self.o_proj_fp8 = (self.o_groups // self.tp_world_size_) == 1 + if self.o_proj_fp8: + self.wo_a_ = ROWMMWeight( + in_dim=per_group_in, + out_dims=[self.o_groups * self.o_lora_rank], + weight_names=f"{p}.wo_a.weight", + data_type=self.data_type_, + quant_method=self.get_quant_method("wo_a"), + ) + else: + self.wo_a_ = ROWBMMWeight( + dim0=self.o_groups, + dim1=per_group_in, + dim2=self.o_lora_rank, + weight_names=f"{p}.wo_a.weight", + data_type=self.data_type_, + quant_method=None, + ) self.wo_b_ = COLMMWeight( in_dim=self.o_groups * self.o_lora_rank, out_dims=[self.hidden], @@ -161,19 +169,12 @@ def _init_indexer(self): tp_world_size=1, ) coff = 2 # indexer compressor always uses ratio 4 (overlap) - self.idx_cmp_wkv_ = ROWMMWeight( - in_dim=self.hidden, - out_dims=[coff * self.index_head_dim], - weight_names=f"{p}.compressor.wkv.weight", - data_type=self.data_type_, - quant_method=None, - tp_rank=0, - tp_world_size=1, - ) - self.idx_cmp_wgate_ = ROWMMWeight( + # wkv/wgate share the same input -> one fused bf16 GEMM producing the [kv | gate] layout + # directly (same as the attention compressor_wkv_gate_). + self.idx_cmp_wkv_gate_ = ROWMMWeight( in_dim=self.hidden, - out_dims=[coff * self.index_head_dim], - weight_names=f"{p}.compressor.wgate.weight", + out_dims=[coff * self.index_head_dim, coff * self.index_head_dim], + weight_names=[f"{p}.compressor.wkv.weight", f"{p}.compressor.wgate.weight"], data_type=self.data_type_, quant_method=None, tp_rank=0, @@ -321,10 +322,13 @@ def _dequant_in_place(self, weights): else: weights[k] = dequant_fp8_block_to_bf16(weights[k], weights[scale_k]).to(self.data_type_) del weights[scale_k] - # grouped-O: reshape [groups*o_lora, in] -> [groups, in, o_lora] for the batched matmul - woa = f"{self.prefix}.attn.wo_a.weight" - if woa in weights and weights[woa].dim() == 2: - w = weights[woa] - per_group_in = self.n_heads * self.head_dim // self.o_groups - weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2).contiguous() + # grouped-O (bf16 path only): reshape [groups*o_lora, in] -> [groups, in, o_lora] for the + # batched matmul. The fp8 path keeps wo_a as a plain [groups*o_lora, in] fp8 GEMM weight + # (its `.scale` is renamed to `.weight_scale_inv` by the loop above, not dequantized). + if not self.o_proj_fp8: + woa = f"{self.prefix}.attn.wo_a.weight" + if woa in weights and weights[woa].dim() == 2: + w = weights[woa] + per_group_in = self.n_heads * self.head_dim // self.o_groups + weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2).contiguous() return From 0d52fde7b9c98c82d9af6cac3364fa47daced76a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 28 Jun 2026 12:40:08 +0000 Subject: [PATCH 051/214] reduce alloc fragment --- .../attention/nsa/fp8_flashmla_sparse.py | 120 ++++++++++++++++-- lightllm/models/deepseek_v4/infer_struct.py | 35 ++++- .../layer_infer/transformer_layer_infer.py | 48 ++++--- lightllm/models/deepseek_v4/model.py | 2 + .../build_compress_index_dsv4.py | 17 +-- .../triton_kernel/build_swa_index_dsv4.py | 19 ++- lightllm/models/deepseek_v4/workspace.py | 51 ++++++++ 7 files changed, 230 insertions(+), 62 deletions(-) create mode 100644 lightllm/models/deepseek_v4/workspace.py diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index 115d5d1934..add32918c3 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -1,25 +1,44 @@ import dataclasses +import inspect import torch from typing import TYPE_CHECKING, Tuple from ..base_att import AttControl, BaseAttBackend, BaseDecodeAttState, BasePrefillAttState from lightllm.utils.dist_utils import get_current_device_id +from lightllm.utils.log_utils import init_logger if TYPE_CHECKING: from lightllm.common.basemodel.infer_struct import InferStateInfo +logger = init_logger(__name__) # this flash_mla extra-cache fork only instantiates h_q in {64, 128}; pad TP-split q heads up # to the nearest supported count (zero heads are discarded from the output slice). FLASHMLA_SUPPORTED_HEADS = (64, 128) -def _pad_q_heads(q_4d: torch.Tensor, attn_sink: torch.Tensor): +def _target_q_heads(h_q: int) -> int: + target = next((h for h in FLASHMLA_SUPPORTED_HEADS if h >= h_q), None) + assert target is not None, f"num q heads {h_q} exceeds flash_mla support {FLASHMLA_SUPPORTED_HEADS}" + return target + + +def _pad_q_heads( + q_4d: torch.Tensor, + attn_sink: torch.Tensor, + q_out: torch.Tensor = None, + sink_out: torch.Tensor = None, +): h_q = q_4d.shape[2] if h_q in FLASHMLA_SUPPORTED_HEADS: return q_4d, attn_sink, h_q - target = next((h for h in FLASHMLA_SUPPORTED_HEADS if h >= h_q), None) - assert target is not None, f"num q heads {h_q} exceeds flash_mla support {FLASHMLA_SUPPORTED_HEADS}" + target = _target_q_heads(h_q) + if q_out is not None: + q_out[:, :, :h_q, :].copy_(q_4d) + q_out[:, :, h_q:target, :].zero_() + sink_out[:h_q].copy_(attn_sink) + sink_out[h_q:target].zero_() + return q_out, sink_out[:target], h_q q_pad = torch.nn.functional.pad(q_4d, (0, 0, 0, target - h_q)) sink_pad = torch.nn.functional.pad(attn_sink, (0, target - h_q)) return q_pad, sink_pad, h_q @@ -71,6 +90,61 @@ def __init__(self, model): torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) for _ in range(2) ] + self.prefill_flash_mla, self.prefill_flash_mla_supports_out = self._load_prefill_flash_mla() + self.prefill_q_workspace = None + self.prefill_out_workspace = None + self.prefill_real_out_workspace = None + self.prefill_sink_workspace = None + self.prefill_workspace_shape = None + self.prefill_real_out_shape = None + self.prefill_workspace_token_capacity = int(model.batch_max_tokens or 0) + if self.prefill_flash_mla_supports_out: + logger.info("DSV4 FlashMLA prefill uses vLLM out= workspace path") + else: + logger.warning("DSV4 FlashMLA prefill out= path unavailable; falling back to allocating FlashMLA output") + + def _load_prefill_flash_mla(self): + try: + from vllm.v1.attention.ops import flashmla as flash_mla + + sig = inspect.signature(flash_mla.flash_mla_with_kvcache) + if "out" in sig.parameters: + return flash_mla, True + except Exception: + pass + + import flash_mla + + return flash_mla, "out" in inspect.signature(flash_mla.flash_mla_with_kvcache).parameters + + def _ensure_prefill_workspace(self, token_num: int, head_num: int, target_heads: int, head_dim: int, dtype, device): + capacity = max(token_num, self.prefill_workspace_token_capacity) + workspace_shape = (capacity, 1, target_heads, head_dim) + if ( + self.prefill_workspace_shape != (target_heads, head_dim, dtype, device) + or self.prefill_q_workspace is None + or self.prefill_q_workspace.shape[0] < capacity + ): + self.prefill_q_workspace = torch.empty(workspace_shape, dtype=dtype, device=device) + self.prefill_out_workspace = torch.empty(workspace_shape, dtype=dtype, device=device) + self.prefill_sink_workspace = torch.empty((target_heads,), dtype=torch.float32, device=device) + self.prefill_workspace_shape = (target_heads, head_dim, dtype, device) + + real_out_shape = (capacity, head_num, head_dim) + if ( + self.prefill_real_out_shape != (head_num, head_dim, dtype, device) + or self.prefill_real_out_workspace is None + or self.prefill_real_out_workspace.shape[0] < capacity + ): + self.prefill_real_out_workspace = torch.empty(real_out_shape, dtype=dtype, device=device) + self.prefill_real_out_shape = (head_num, head_dim, dtype, device) + + return ( + self.prefill_q_workspace[:token_num], + self.prefill_out_workspace[:token_num], + self.prefill_real_out_workspace[:token_num], + self.prefill_sink_workspace, + ) def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparsePrefillAttState": return NsaFlashMlaFp8SparsePrefillAttState(backend=self, infer_state=infer_state) @@ -120,9 +194,7 @@ def ensure_nsa_ks_ke(self): def _get_flashmla_sched_meta(self, compress_ratio: int): sched_meta = self.flashmla_sched_meta.get(compress_ratio) if sched_meta is None: - import flash_mla - - sched_meta = flash_mla.get_mla_metadata()[0] + sched_meta = self.backend.prefill_flash_mla.get_mla_metadata()[0] self.flashmla_sched_meta[compress_ratio] = sched_meta return sched_meta @@ -133,6 +205,7 @@ def prefill_att( v: torch.Tensor, att_control: AttControl = AttControl(), alloc_func=torch.empty, + out: torch.Tensor = None, ) -> torch.Tensor: assert att_control.nsa_prefill, "nsa_prefill must be True for NSA prefill attention" assert att_control.nsa_prefill_dict is not None, "nsa_prefill_dict is required" @@ -141,6 +214,7 @@ def prefill_att( q=q, packed_kv=k, nsa_dict=att_control.nsa_prefill_dict, + out=out, ) return self._nsa_prefill_att(q=q, packed_kv=k, att_control=att_control) @@ -188,10 +262,12 @@ def _nsa_prefill_att( ) return mla_out - def _flashmla_kvcache_prefill_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: + def _flashmla_kvcache_prefill_att( + self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict, out: torch.Tensor = None + ) -> torch.Tensor: attn_sink = nsa_dict["attn_sink"] metadata = _metadata_from_dict(self.infer_state, nsa_dict) - return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) + return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict, out=out) def _flashmla_kvcache_att( self, @@ -200,16 +276,27 @@ def _flashmla_kvcache_att( metadata: _Dsv4Metadata, attn_sink: torch.Tensor, nsa_dict: dict, + out: torch.Tensor = None, ) -> torch.Tensor: - import flash_mla from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE q_4d = q.unsqueeze(1).contiguous() - q_4d, attn_sink, num_real_heads = _pad_q_heads(q_4d, attn_sink) + num_real_heads = q_4d.shape[2] + target_heads = _target_q_heads(num_real_heads) + q_workspace, full_out_workspace, real_out_workspace, sink_workspace = self.backend._ensure_prefill_workspace( + q_4d.shape[0], + num_real_heads, + target_heads, + q_4d.shape[-1], + q_4d.dtype, + q_4d.device, + ) + q_for_flash, sink_for_flash, num_real_heads = _pad_q_heads(q_4d, attn_sink, q_workspace, sink_workspace) k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) sched_meta = self._get_flashmla_sched_meta(nsa_dict["compress_ratio"]) - out, _ = flash_mla.flash_mla_with_kvcache( - q=q_4d, + flash_mla = self.backend.prefill_flash_mla + kwargs = dict( + q=q_for_flash, k_cache=k_cache, block_table=None, cache_seqlens=None, @@ -220,13 +307,18 @@ def _flashmla_kvcache_att( causal=False, is_fp8_kvcache=True, indices=metadata.swa_indices, - attn_sink=attn_sink, + attn_sink=sink_for_flash, topk_length=metadata.swa_lengths, extra_k_cache=metadata.extra_cache, extra_indices_in_kvcache=metadata.extra_indices, extra_topk_length=metadata.extra_lengths, ) - return out[:, 0, :num_real_heads].contiguous() + if self.backend.prefill_flash_mla_supports_out: + kwargs["out"] = full_out_workspace + full_out, _ = flash_mla.flash_mla_with_kvcache(**kwargs) + real_out = out if out is not None else real_out_workspace + real_out.copy_(full_out[:, 0, :num_real_heads, :]) + return real_out @dataclasses.dataclass diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index f37620ddff..e14e84c968 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -23,6 +23,9 @@ def __init__(self): self.dsv4_sparse_req_idx = None self.dsv4_swa_indices = None self.dsv4_swa_lengths = None + self.dsv4_c128_indices = None + self.dsv4_c128_lengths = None + self.dsv4_workspace = None # token -> batch-position map for the compressor; built per prefill forward in init_some_extra_state. self._dsv4_token_to_batch_idx = None # lazily-built (first c4 layer) cache of layer-independent paged-c4 metadata; reused by the @@ -30,6 +33,15 @@ def __init__(self): # ignores it -- it's a capture-time wiring of layer0->others, not a staged graph input. self._c4_paged_meta = None + def _dsv4_index_max_kv_seq_len(self, model): + if ( + not self.is_prefill + and model.graph is not None + and model.graph.can_run(self.batch_size, self.max_kv_seq_len) + ): + return model.graph.graph_max_len_in_batch + return self.max_kv_seq_len + def init_some_extra_state(self, model): self._c4_paged_meta = None # reset per forward before any c4 layer runs super().init_some_extra_state(model) # sets position_ids, b_q_seq_len, b_q_start_loc (prefill) @@ -50,17 +62,32 @@ def init_some_extra_state(self, model): else: self.dsv4_sparse_req_idx = self.b_req_idx self._dsv4_token_to_batch_idx = None - # Sliding-window indices: layer-independent (full_to_swa is global, window is const), so build - # once here via one fused kernel instead of recomputing per layer. const [T, window] shape is - # cuda-graph-safe (no max_kv_seq_len dependence) and auto-staged by copy_for_cuda_graph. + # Sliding-window indices are layer-independent, so build them once into the model workspace. from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_swa_index + workspace = model.dsv4_workspace + self.dsv4_workspace = workspace + self.dsv4_swa_indices, self.dsv4_swa_lengths = workspace.swa(self.microbatch_index, pos.numel()) self.dsv4_swa_indices, self.dsv4_swa_lengths = build_swa_index( req_idx=self.dsv4_sparse_req_idx, positions=self.position_ids, req_to_token_indexs=self.req_manager.req_to_token_indexs, full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, - window=int(self.mem_manager.sliding_window), + swa_index=self.dsv4_swa_indices, + swa_length=self.dsv4_swa_lengths, + ) + from lightllm.models.deepseek_v4.triton_kernel.build_compress_index_dsv4 import build_compress_index + + cap = workspace.compress_cap(self._dsv4_index_max_kv_seq_len(model), 128) + self.dsv4_c128_indices, self.dsv4_c128_lengths = workspace.c128(self.microbatch_index, pos.numel(), cap) + build_compress_index( + self.dsv4_sparse_req_idx, + self.position_ids, + self.req_manager.req_to_token_indexs, + self.mem_manager.full_to_c128_indexs, + 128, + self.dsv4_c128_indices, + self.dsv4_c128_lengths, ) # prefill-cudagraph 桶填充的 HOLD 尾请求的 q 行数。其注意力读 HOLD 槽位(内容被并发写 # 竞争,每轮不同),输出必须清零,否则 pad 行 hidden 不确定 -> MoE 路由抖动 -> 共享 expert diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index d8b4955d6f..a5f37936cd 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -123,7 +123,8 @@ def context_forward( x = self.context_attention_forward(x, infer_state, layer_weight) x, residual, post_mix, res_mix = self._hc_ffn_in(x, residual, post_mix, res_mix, layer_weight) x = self._ffn(x, infer_state, layer_weight) - return self._hc_ffn_out(x, residual, post_mix, res_mix) + out = self._hc_ffn_out(x, residual, post_mix, res_mix) + return out def token_forward( self, input_embdings, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight @@ -220,9 +221,10 @@ def _context_attention_wrapper_run( _o = tensor_to_no_ref_tensor(o) def att_func(new_infer_state: DeepseekV4InferStateInfo): - tmp_o = self._context_attention_kernel(_q, _q_lora, _x, new_infer_state, layer_weight) + tmp_o = self._context_attention_kernel(_q, _q_lora, _x, new_infer_state, layer_weight, out=_o) assert tmp_o.shape == _o.shape - _o.copy_(tmp_o) + if tmp_o.data_ptr() != _o.data_ptr(): + _o.copy_(tmp_o) return infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) @@ -265,7 +267,13 @@ def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, return self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) def _context_attention_kernel( - self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + self, + q, + q_lora, + x, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + out=None, ): meta = self._compress_and_index(q_lora, x, infer_state, layer_weight) att_control = AttControl( @@ -285,6 +293,7 @@ def _context_attention_kernel( k=infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_), v=None, att_control=att_control, + out=out, ) pad_q_len = getattr(infer_state, "_dsv4_prefill_pad_q_len", 0) if pad_q_len: @@ -495,7 +504,7 @@ class DeepseekV4IndexInfer: ragged gather of the compressed c4 keys, deep_gemm.fp8_mqa_logits, then topk -- adapted for the replicated indexer (no gather-q/all_reduce), the c4-compressed entry space, and topk-512 (no inheritance only because of those data-shape differences). swa metadata is precomputed in - init_some_extra_state; this class owns the c4/c128 entry gather (build_compress_index) AND the c4 + init_some_extra_state; this class owns the c4 entry gather (build_compress_index) AND the c4 Lightning-Indexer scoring (gather + deep_gemm.fp8_mqa_logits + topk). Holds only static per-layer config; all per-request data flows in via args. Invoke from _context/_token_attention_kernel (after compressor.fused_compress, before *_att) so the c4 scorer/topk keep the same cuda-graph @@ -558,11 +567,10 @@ def build_metadata( ): """Return the final flash_mla index tensors for this layer's compress variant. swa indices and the per-token req_idx are layer-independent and precomputed once in init_some_extra_state - (read here); only the c4 scorer / c128 gather is per-layer. The backend pairs these with the + (read here); only the c4 scorer is per-layer. The backend pairs these with the (data-independent, layer-keyed) fp8 cache-byte views it owns.""" swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) swa_lengths = infer_state.dsv4_swa_lengths - req_idx = infer_state.dsv4_sparse_req_idx positions = infer_state.position_ids extra_indices = extra_lengths = None if self.compress_ratio == 4: @@ -571,7 +579,8 @@ def build_metadata( ) extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: - extra_indices, extra_lengths = self._c128_indices(infer_state, req_idx, positions) + extra_indices = infer_state.dsv4_c128_indices.unsqueeze(1) + extra_lengths = infer_state.dsv4_c128_lengths return { "swa_indices": swa_indices, "swa_lengths": swa_lengths, @@ -579,20 +588,6 @@ def build_metadata( "extra_lengths": extra_lengths, } - def _c128_indices(self, infer_state: DeepseekV4InferStateInfo, req_idx, positions): - from ..triton_kernel.build_compress_index_dsv4 import build_compress_index - - cap = ((max(1, int(infer_state.max_kv_seq_len) // 128) + 63) // 64) * 64 - indices, lengths = build_compress_index( - req_idx, - positions, - infer_state.req_manager.req_to_token_indexs, - infer_state.mem_manager.full_to_c128_indexs, - ratio=128, - cap=cap, - ) - return indices.unsqueeze(1), lengths - def _indexer_q_weight( self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True ): @@ -627,6 +622,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, then masked topk-512 -> c4 slots. Fixed shapes (c4_cap pinned per graph bucket) keep the decode cuda graph capturable.""" mem_manager = infer_state.mem_manager + workspace = infer_state.dsv4_workspace index_topk = self.index_topk max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) c4_cap = ((max_entries + 63) // 64) * 64 @@ -637,13 +633,15 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, if max_entries <= index_topk: from ..triton_kernel.build_compress_index_dsv4 import build_compress_index + slots, lengths = workspace.c4(infer_state.microbatch_index, positions.shape[0], c4_cap) slots, lengths = build_compress_index( infer_state.dsv4_sparse_req_idx, positions, infer_state.req_manager.req_to_token_indexs, mem_manager.full_to_c4_indexs, - ratio=4, - cap=c4_cap, + 4, + slots, + lengths, ) return slots.unsqueeze(1), lengths @@ -713,7 +711,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, c4_cap, False, ) - top_slots = torch.empty((idx_q_fp8.shape[0], index_topk), dtype=torch.int32, device=device) + top_slots, _ = workspace.c4(infer_state.microbatch_index, idx_q_fp8.shape[0], index_topk) topk_transform_512( logits, valid_len, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 5353dfcade..edb2ce15db 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -25,6 +25,7 @@ ) from lightllm.common.basemodel.attention import get_nsa_prefill_att_backend_class, get_nsa_decode_att_backend_class from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.deepseek_v4.workspace import DeepseekV4Workspace from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, linear_ramp_mask, @@ -113,6 +114,7 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() + self.dsv4_workspace = DeepseekV4Workspace(self) if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1": prefill_aux_stream = torch.cuda.Stream() for layer in self.layers_infer: diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py index b09192498d..dae5121ab9 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py @@ -11,6 +11,7 @@ def _build_compress_index_kernel( req_to_token_stride0, full_to_c_ptr, index_ptr, + index_stride0, length_ptr, cap, RATIO: tl.constexpr, @@ -30,7 +31,7 @@ def _build_compress_index_kernel( safe_pos = tl.where(valid, end_pos, 0) full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_pos, mask=valid, other=0).to(tl.int64) c_slot = tl.load(full_to_c_ptr + full_slot, mask=valid, other=-1).to(tl.int32) - tl.store(index_ptr + t * cap + e, c_slot, mask=e_mask) + tl.store(index_ptr + t * index_stride0 + e, c_slot, mask=e_mask) if eb == 0: tl.store(length_ptr + t, tl.maximum(raw_len, 1).to(tl.int32)) @@ -42,22 +43,21 @@ def build_compress_index( req_to_token_indexs: torch.Tensor, full_to_c_indexs: torch.Tensor, ratio: int, - cap: int, + index: torch.Tensor, + length: torch.Tensor, ): """Fused two-level group-end gather for the c4/c128 compressed-entry index tables. For token t (at request `req_idx[t]`, absolute `positions[t]`) and compressed entry e: slot[t, e] = full_to_c[ req_to_token[req, e*ratio + (ratio-1)] ] (the group-end token's full slot) with slot = -1 where e >= (pos+1)//ratio (beyond the causal compressed length) or where the - full->c map is unset. Returns (index [T, cap] int32, length [T] int32 = clamp((pos+1)//ratio, 1)). + full->c map is unset. Writes index [T, cap] and length [T] = clamp((pos+1)//ratio, 1). - Replaces the eager _gather_compress_slots/_c128/c4-causal torch chain. `cap` must be a multiple of - 64 (FlashMLA topk alignment); the tiled grid (T, ceil(cap/BLOCK_E)) scales to 1M-context caps. - cuda-graph-safe: cap is fixed per graph bucket, shapes static. + Replaces the eager _gather_compress_slots/_c128/c4-causal torch chain. The caller owns the + output storage, so this wrapper does not allocate on the hot path. """ T = positions.shape[0] - index = torch.empty((T, cap), dtype=torch.int32, device=positions.device) - length = torch.empty((T,), dtype=torch.int32, device=positions.device) + cap = index.shape[1] if T == 0: return index, length BLOCK_E = 256 @@ -69,6 +69,7 @@ def build_compress_index( req_to_token_indexs.stride(0), full_to_c_indexs, index, + index.stride(0), length, cap, RATIO=ratio, diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py index e5b7b80bb4..a1ef2d5be1 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py @@ -11,6 +11,7 @@ def _build_swa_index_kernel( req_to_token_stride0, full_to_swa_ptr, swa_index_ptr, + swa_index_stride0, swa_length_ptr, WINDOW: tl.constexpr, BLOCK_W: tl.constexpr, @@ -28,7 +29,7 @@ def _build_swa_index_kernel( full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_offset, mask=valid, other=0).to(tl.int64) swa_slot = tl.load(full_to_swa_ptr + full_slot, mask=valid, other=-1) out = tl.where(valid, swa_slot, -1).to(tl.int32) - tl.store(swa_index_ptr + token_idx * WINDOW + w, out, mask=w_mask) + tl.store(swa_index_ptr + token_idx * swa_index_stride0 + w, out, mask=w_mask) length = tl.minimum(tl.maximum(pos + 1, 1), WINDOW).to(tl.int32) tl.store(swa_length_ptr + token_idx, length) @@ -39,7 +40,8 @@ def build_swa_index( positions: torch.Tensor, req_to_token_indexs: torch.Tensor, full_to_swa_indexs: torch.Tensor, - window: int, + swa_index: torch.Tensor, + swa_length: torch.Tensor, ): """Per-token sliding-window FlashMLA index table, built ONCE per forward (layer-independent: full_to_swa is a single global map and the window is a model constant, so every layer's swa @@ -47,17 +49,11 @@ def build_swa_index( (req_idx, position) gather the last `window` tokens' full slots via req_to_token, then map full -> swa; out-of-range positions store -1. - Returns (swa_index [T, window] int32, swa_length [T] int32). `window` is 128 (a multiple of the - FlashMLA 64 alignment) so no extra pad is needed; the reader adds the s_q axis via unsqueeze(1). - Const output shape (no max_kv_seq_len dependence) makes this cuda-graph-safe to stage from - init_some_extra_state via copy_for_cuda_graph. + Writes (swa_index [T, window] int32, swa_length [T] int32). The caller owns the output storage; + the reader adds the s_q axis via unsqueeze(1). """ - # window must stay 64-aligned: the output is the FlashMLA `indices` tensor directly (no separate - # _pad_last_dim), and the extra-cache fork requires the topk dim to be a multiple of 64. - assert window % 64 == 0, f"DeepSeek-V4 sliding_window must be a multiple of 64 for FlashMLA, got {window}" T = positions.shape[0] - swa_index = torch.empty((T, window), dtype=torch.int32, device=positions.device) - swa_length = torch.empty((T,), dtype=torch.int32, device=positions.device) + window = swa_index.shape[1] if T == 0: return swa_index, swa_length _build_swa_index_kernel[(T,)]( @@ -67,6 +63,7 @@ def build_swa_index( req_to_token_indexs.stride(0), full_to_swa_indexs, swa_index, + swa_index.stride(0), swa_length, WINDOW=window, BLOCK_W=triton.next_power_of_2(window), diff --git a/lightllm/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py new file mode 100644 index 0000000000..562a53bbf3 --- /dev/null +++ b/lightllm/models/deepseek_v4/workspace.py @@ -0,0 +1,51 @@ +import torch +from lightllm.utils.envs_utils import get_env_start_args + + +class DeepseekV4Workspace: + def __init__(self, model): + self.token_capacity = int(model.batch_max_tokens) + self.sliding_window = int(model.config["sliding_window"]) + self.index_topk = int(model.config["index_topk"]) + self.c128_cap = self.compress_cap(model.max_seq_length, 128) + args = get_env_start_args() + overlap = args.enable_decode_microbatch_overlap or args.enable_prefill_microbatch_overlap + self.microbatch_count = 1 + int(overlap) + + self.swa_indices = self._alloc(self.sliding_window) + self.swa_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") + self.c4_indices = self._alloc(self.index_topk) + self.c4_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") + self.c128_indices = self._alloc(self.c128_cap) + self.c128_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") + + @staticmethod + def compress_cap(max_kv_seq_len: int, ratio: int) -> int: + entries = max(1, int(max_kv_seq_len) // ratio) + return ((entries + 63) // 64) * 64 + + def _alloc(self, width: int) -> torch.Tensor: + return torch.empty((self.microbatch_count, self.token_capacity * width), dtype=torch.int32, device="cuda") + + @staticmethod + def _view(buffer: torch.Tensor, token_num: int, width: int) -> torch.Tensor: + return torch.as_strided(buffer, (token_num, width), (width, 1)) + + def swa(self, microbatch_index: int, token_num: int): + return ( + self._view(self.swa_indices[microbatch_index], token_num, self.sliding_window), + self.swa_lengths[microbatch_index, :token_num], + ) + + def c4(self, microbatch_index: int, token_num: int, width: int): + assert width <= self.index_topk, f"c4 width {width} exceeds allocated {self.index_topk}" + return ( + self._view(self.c4_indices[microbatch_index], token_num, width), + self.c4_lengths[microbatch_index, :token_num], + ) + + def c128(self, microbatch_index: int, token_num: int, width: int): + return ( + self._view(self.c128_indices[microbatch_index], token_num, width), + self.c128_lengths[microbatch_index, :token_num], + ) From c711fe96eb3bbe49681268c32b32a7dfcc0560c2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 28 Jun 2026 13:53:30 +0000 Subject: [PATCH 052/214] set _C4_PREFILL_LOGITS_BUDGET_BYTES reduce max memory usage --- .../layer_infer/transformer_layer_infer.py | 57 ++++++++++++++++++- 1 file changed, 55 insertions(+), 2 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index a5f37936cd..664bc23576 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -18,6 +18,9 @@ from lightllm.third_party.sglang_jit.dsv4 import topk_transform_512 +_C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 + + class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): def __init__(self, layer_num, network_config): TransformerLayerInferTpl.__init__(self, layer_num, network_config) @@ -701,6 +704,58 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, 1, self.index_head_dim + 4, ) + top_slots, _ = workspace.c4(infer_state.microbatch_index, idx_q_fp8.shape[0], index_topk) + if infer_state.is_prefill: + rows_per_chunk = max(1, _C4_PREFILL_LOGITS_BUDGET_BYTES // (c4_cap * 4)) + if idx_q_fp8.shape[0] > rows_per_chunk: + for start in range(0, idx_q_fp8.shape[0], rows_per_chunk): + end = min(start + rows_per_chunk, idx_q_fp8.shape[0]) + chunk_ctx_lens = ctx_lens[start:end] + self._c4_score_topk( + idx_q_fp8[start:end], + kv_cache, + weights[start:end], + chunk_ctx_lens, + row_page_table[start:end], + deep_gemm.get_paged_mqa_logits_metadata( + chunk_ctx_lens, + page_size, + deep_gemm.get_num_sms(), + ), + c4_cap, + valid_len[start:end], + top_slots[start:end], + page_size, + ) + return top_slots.unsqueeze(1), topk_lengths + + self._c4_score_topk( + idx_q_fp8, + kv_cache, + weights, + ctx_lens, + row_page_table, + metadata, + c4_cap, + valid_len, + top_slots, + page_size, + ) + return top_slots.unsqueeze(1), topk_lengths + + @staticmethod + def _c4_score_topk( + idx_q_fp8, + kv_cache, + weights, + ctx_lens, + row_page_table, + metadata, + c4_cap, + valid_len, + top_slots, + page_size, + ): logits = deep_gemm.fp8_paged_mqa_logits( idx_q_fp8.unsqueeze(1), kv_cache, @@ -711,7 +766,6 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, c4_cap, False, ) - top_slots, _ = workspace.c4(infer_state.microbatch_index, idx_q_fp8.shape[0], index_topk) topk_transform_512( logits, valid_len, @@ -719,4 +773,3 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, top_slots, page_size, ) - return top_slots.unsqueeze(1), topk_lengths From fced38bcfdf4e21936c58c3b893ad854322405fe Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 1 Jul 2026 08:54:23 +0000 Subject: [PATCH 053/214] fix(deepseek-v4): align fp8 serving numerics with reference - Apply DeepSeek YaRN correction to both sliding-window and compressed RoPE tables. The two paths use different RoPE bases, but both need the same YaRN correction range, matching SGLang/vLLM behavior. - Keep DeepSeek-V4 wo_a in BF16 on Hopper and only enable the FP8 single-group O-proj path on SM100+. The checkpoint stores wo_a as BF16, and SGLang also leaves this projection BF16 on Hopper; forcing FP8 there changes attention numerics. - Add an out_dtype path to MMWeightTpl/NoQuantization and use it for DS4 BF16 GEMMs that must produce FP32 outputs. This keeps the call style inside the LightLLM weight abstraction while matching SGLang's linear_bf16_fp32 behavior for router logits and compressor kv_score. - Apply swiglu_limit to the shared expert MLP path as well as routed experts. DS4 shared experts use the same clamped SiLU-and-mul behavior in the reference implementation, so the shared/routed sum now matches the expected MLP math. --- .../meta_weights/mm_weight/mm_weight.py | 17 ++++++++- lightllm/common/quantization/no_quant.py | 9 ++++- .../common/quantization/quantize_method.py | 1 + .../layer_infer/transformer_layer_infer.py | 28 +++++++++++---- .../layer_weights/transformer_layer_weight.py | 19 ++++++---- lightllm/models/deepseek_v4/model.py | 35 +++++++++++-------- 6 files changed, 78 insertions(+), 31 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_weight.py index 5021699143..1b966c5738 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/mm_weight/mm_weight.py @@ -54,8 +54,23 @@ def __init__( self.gen_weight_quant_param_names() def mm( - self, input_tensor: torch.Tensor, out: Optional[torch.Tensor] = None, use_custom_tensor_mananger: bool = True + self, + input_tensor: torch.Tensor, + out: Optional[torch.Tensor] = None, + use_custom_tensor_mananger: bool = True, + out_dtype: Optional[torch.dtype] = None, ) -> torch.Tensor: + if out_dtype is not None and not isinstance(self.quant_method, NoQuantization): + raise NotImplementedError(f"out_dtype is not supported for quant method {self.quant_method.method_name}") + if out_dtype is not None: + return self.quant_method.apply( + input_tensor, + self.mm_param, + out, + use_custom_tensor_mananger=use_custom_tensor_mananger, + bias=self.bias, + out_dtype=out_dtype, + ) return self.quant_method.apply( input_tensor, self.mm_param, out, use_custom_tensor_mananger=use_custom_tensor_mananger, bias=self.bias ) diff --git a/lightllm/common/quantization/no_quant.py b/lightllm/common/quantization/no_quant.py index fa926ad6f0..21e7101d4d 100644 --- a/lightllm/common/quantization/no_quant.py +++ b/lightllm/common/quantization/no_quant.py @@ -18,19 +18,26 @@ def apply( workspace: Optional[torch.Tensor] = None, use_custom_tensor_mananger: bool = True, bias: Optional[torch.Tensor] = None, + out_dtype: Optional[torch.dtype] = None, ) -> torch.Tensor: from lightllm.common.basemodel.layer_infer.cache_tensor_manager import g_cache_manager weight = weight_pack.weight.t() + if out_dtype is not None and bias is not None: + raise NotImplementedError("out_dtype is only supported for bias-free no-quant mm") if out is None: shape = (input_tensor.shape[0], weight.shape[1]) - dtype = input_tensor.dtype + dtype = out_dtype if out_dtype is not None else input_tensor.dtype device = input_tensor.device if use_custom_tensor_mananger: out = g_cache_manager.alloc_tensor(shape, dtype, device=device) else: out = torch.empty(shape, dtype=dtype, device=device) + elif out_dtype is not None and out.dtype != out_dtype: + raise ValueError(f"out dtype {out.dtype} does not match requested out_dtype {out_dtype}") if bias is None: + if out_dtype is not None: + return torch.mm(input_tensor, weight, out=out, out_dtype=out_dtype) return torch.mm(input_tensor, weight, out=out) return torch.addmm(bias, input_tensor, weight, out=out) diff --git a/lightllm/common/quantization/quantize_method.py b/lightllm/common/quantization/quantize_method.py index 95d8d806f9..d3f251ec84 100644 --- a/lightllm/common/quantization/quantize_method.py +++ b/lightllm/common/quantization/quantize_method.py @@ -55,6 +55,7 @@ def apply( workspace: Optional[torch.Tensor] = None, use_custom_tensor_mananger: bool = True, bias: Optional[torch.Tensor] = None, + out_dtype: Optional[torch.dtype] = None, ) -> torch.Tensor: pass diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 664bc23576..d67ad11165 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -2,6 +2,7 @@ import torch.distributed as dist from lightllm.common.basemodel import TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl +from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.distributed.communication_op import all_reduce from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight @@ -47,7 +48,7 @@ def __init__(self, layer_num, network_config): self.sin_compress_table = None self.num_experts_per_tok = network_config["num_experts_per_tok"] self.routed_scaling_factor = network_config["routed_scaling_factor"] - self.swiglu_limit = network_config["swiglu_limit"] + self.swiglu_limit = float(network_config["swiglu_limit"]) self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) self.tp_q_head_num_ = self.num_heads // self.tp_world_size_ self.tp_groups = self.o_groups // self.tp_world_size_ @@ -347,17 +348,28 @@ def _routed_experts(self, x, weights, indices, layer_weight: DeepseekV4Transform clamp_limit=float(self.swiglu_limit), ) + def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): + input = input.view(-1, self.embed_dim_) + gate_up = layer_weight.gate_up_proj.mm(input) + shared = self.alloc_tensor((input.size(0), gate_up.size(1) // 2), input.dtype) + silu_and_mul_fwd(gate_up, shared, limit=self.swiglu_limit) + input = None + gate_up = None + out = layer_weight.down_proj.mm(shared) + shared = None + return out + def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): x = x.view(-1, self.embed_dim_) if not self.enable_ep_moe: x = self._tpsp_allgather(input=x, infer_state=infer_state) - logits = layer_weight.gate_weight_.mm(x).float().contiguous() + logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) weights, indices = self._select_experts(logits, infer_state, layer_weight) # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 - # 复用 Llama 的 _ffn_tp: fused gate_up matmul + silu_and_mul triton kernel,无 swiglu clamp, - # 对齐参考 DeepseekV4MLP(=LlamaMLP)。swiglu_limit clamp 只属于 routed 专家 (见 _routed_experts)。 + # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's + # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) routed = self._routed_experts(x, weights, indices, layer_weight) if self.enable_ep_moe: @@ -451,11 +463,13 @@ def prepare_states( # fused wkv/wgate GEMM -> [T, 2*coff*idx_hd] in the [kv | score] layout directly # (same as the attention compressor_wkv_gate_). self._metadata.kv_score = layer_weight.idx_cmp_wkv_gate_.mm( - x, use_custom_tensor_mananger=use_custom_tensor_manager - ).float() + x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 + ) ape = layer_weight.idx_cmp_ape_.weight else: - self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm(x).float() + self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm( + x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 + ) ape = layer_weight.compressor_ape_.weight prepare_partial_states( kv_score=self._metadata.kv_score, diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 5ecce0a763..072591a552 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -89,12 +89,13 @@ def _init_qkvo(self): # grouped low-rank output projection (wo_a per-group [in, o_lora], wo_b row-parallel # [groups*o_lora -> hidden]). per_group_in = self.n_heads * self.head_dim // self.o_groups - # When o_groups == tp_world_size (e.g. the daily tp8 config) each rank owns exactly ONE - # group, so the grouped O-proj collapses to a single GEMM -> run it in fp8 (deepgemm) - # instead of dequantizing wo_a to bf16. sglang does the same (fp8 wo_a is default-on there). + # When o_groups == tp_world_size (e.g. tp8) each rank owns exactly one group, so the + # grouped O-proj can collapse to a single GEMM. SGLang only enables the fp8 wo_a GEMM + # on Blackwell; on Hopper it keeps BF16, and this checkpoint stores wo_a as BF16. # For >1 group per rank (tp < o_groups) the per-group inputs differ (block-diagonal), so - # keep the bf16 grouped bmm. - self.o_proj_fp8 = (self.o_groups // self.tp_world_size_) == 1 + # keep the BF16 grouped bmm. + major, _ = torch.cuda.get_device_capability() + self.o_proj_fp8 = (self.o_groups // self.tp_world_size_) == 1 and major >= 10 if self.o_proj_fp8: self.wo_a_ = ROWMMWeight( in_dim=per_group_in, @@ -190,8 +191,8 @@ def _init_indexer(self): # ------------------------------------------------------------------ moe def _init_moe(self): p = f"{self.prefix}.ffn" - # Router gate in bf16 (matches the sglang/vLLM DeepSeek references, which run the gate GEMM in - # the model dtype); the bf16 GEMM output is cast back to fp32 in _ffn for topk_hash_softplus_sqrt. + # Router gate weights stay bf16, but DS4 routing consumes fp32 GEMM output + # (SGLang linear_bf16_fp32 / vLLM router_logits_dtype=torch.float32). self.gate_weight_ = ROWMMWeight( in_dim=self.hidden, out_dims=[self.n_routed_experts], @@ -331,4 +332,8 @@ def _dequant_in_place(self, weights): w = weights[woa] per_group_in = self.n_heads * self.head_dim // self.o_groups weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2).contiguous() + # Keep c4 overlap APE in checkpoint layout [4, 2*head_dim]. SGLang reorders it + # because its compressor consumes ape.view(8, head_dim) by window offset. LightLLM + # adds APE into each token's two score halves before compression using position % 4, + # so the raw checkpoint layout is the equivalent representation here. return diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index edb2ce15db..d3f0715b0b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -132,8 +132,8 @@ def _init_to_get_rotary(self): # gemma4 two-variant convention; the fused sglang q kernel consumes them directly, while # _cos_cached_*/_sin_cached_* are .real/.imag views of the same storage for the kv rope, # inverse rope and compressor paths (deepseek2's interleaved triton rotary_emb_fwd). - # Sliding-window layers use base rope_theta (no YaRN); - # compressed (CSA/HCA) layers use compress_rope_theta with configured rope_scaling. + # Sliding-window and compressed layers both use DeepSeek YaRN correction; only the + # RoPE base differs (rope_theta vs compress_rope_theta), matching SGLang/vLLM. # Kept fp32 for accuracy (the apply upcasts anyway). cfg = self.config rs = cfg.get("rope_scaling", {}) or {} @@ -146,22 +146,27 @@ def _init_to_get_rotary(self): freq_exponents = torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim positions = torch.arange(max_seq, dtype=torch.float32, device="cuda") - sliding_freqs = 1.0 / (cfg["rope_theta"] ** freq_exponents) + rope_type = rs.get("rope_type", rs.get("type", "default")) + orig_max = rs.get("original_max_position_embeddings", 0) + + def build_inv_freq(base): + freqs = 1.0 / (base ** freq_exponents) + if rope_type == "yarn" and orig_max > 0: + beta_fast = rs.get("beta_fast", 32) + beta_slow = rs.get("beta_slow", 1) + factor = rs.get("factor", 1) + if factor is None: + factor = cfg.get("max_position_embeddings", max_seq) / orig_max + low, high = find_correction_range(beta_fast, beta_slow, dim, base, orig_max) + smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() + freqs = freqs / factor * (1 - smooth) + freqs * smooth + return freqs + + sliding_freqs = build_inv_freq(cfg["rope_theta"]) f = torch.outer(positions, sliding_freqs) # [max_seq, dim//2] self._freqs_cis_sliding = torch.complex(f.cos(), f.sin()) - compress_freqs = 1.0 / (cfg["compress_rope_theta"] ** freq_exponents) - rope_type = rs.get("rope_type", rs.get("type", "default")) - orig_max = rs.get("original_max_position_embeddings", 0) - if rope_type == "yarn" and orig_max > 0: - beta_fast = rs.get("beta_fast", 32) - beta_slow = rs.get("beta_slow", 1) - factor = rs.get("factor", 1) - if factor is None: - factor = cfg.get("max_position_embeddings", max_seq) / orig_max - low, high = find_correction_range(beta_fast, beta_slow, dim, cfg["compress_rope_theta"], orig_max) - smooth = 1 - linear_ramp_mask(low, high, dim // 2).cuda() - compress_freqs = compress_freqs / factor * (1 - smooth) + compress_freqs * smooth + compress_freqs = build_inv_freq(cfg["compress_rope_theta"]) f = torch.outer(positions, compress_freqs) # [max_seq, dim//2] self._freqs_cis_compress = torch.complex(f.cos(), f.sin()) self._cos_cached_sliding = self._freqs_cis_sliding.real From 17a0d5469541373dace95c789e0d5948b7104879 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 2 Jul 2026 01:52:23 +0000 Subject: [PATCH 054/214] fix stirde bug (if-inverse 73) --- lightllm/models/deepseek_v4/layer_infer/compressor.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 695ac33bc8..e9f2af7385 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -133,8 +133,10 @@ def _fused_compress_norm_rope_insert_kernel( rms_eps, cos_table, cos_stride0, + cos_stride1, sin_table, sin_stride0, + sin_stride1, out_buffer, HEAD_DIM: tl.constexpr, STATE_WIDTH: tl.constexpr, @@ -250,8 +252,8 @@ def _fused_compress_norm_rope_insert_kernel( is_rope_pair = rope_pair_local >= 0 cs_idx = tl.maximum(rope_pair_local, 0) compressed_pos = (position // COMPRESS_RATIO) * COMPRESS_RATIO - cos_v = tl.load(cos_table + compressed_pos * cos_stride0 + cs_idx, mask=is_rope_pair, other=1.0) - sin_v = tl.load(sin_table + compressed_pos * sin_stride0 + cs_idx, mask=is_rope_pair, other=0.0) + cos_v = tl.load(cos_table + compressed_pos * cos_stride0 + cs_idx * cos_stride1, mask=is_rope_pair, other=1.0) + sin_v = tl.load(sin_table + compressed_pos * sin_stride0 + cs_idx * sin_stride1, mask=is_rope_pair, other=0.0) new_even = even * cos_v - odd * sin_v new_odd = odd * cos_v + even * sin_v rotated = tl.interleave(new_even, new_odd) @@ -424,8 +426,10 @@ def fused_compress( eps, cos_table, cos_table.stride(0), + cos_table.stride(1), sin_table, sin_table.stride(0), + sin_table.stride(1), metadata.out_buffer, HEAD_DIM=head_dim, STATE_WIDTH=state_width, From c19b8537134c66040a8dc1468c0b848650155967 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 2 Jul 2026 03:38:21 +0000 Subject: [PATCH 055/214] default topk from huggingface's 50 to -1, if_inverse 70.6 -> 73 --- lightllm/server/api_models.py | 7 ++++--- lightllm/server/core/objs/py_sampling_params.py | 4 ++-- lightllm/server/core/objs/sampling_params.py | 4 ++-- lightllm/utils/config_utils.py | 9 ++++++++- 4 files changed, 16 insertions(+), 8 deletions(-) diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index dcf3d2a0aa..737a08f628 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -4,7 +4,8 @@ from pydantic import BaseModel, Field, field_validator, model_validator from typing import Any, Dict, List, Optional, Union, Literal, ClassVar -from transformers import GenerationConfig + +from lightllm.utils.config_utils import get_generation_config_diff_dict class ImageURL(BaseModel): @@ -160,7 +161,7 @@ class CompletionRequest(BaseModel): def load_generation_cfg(cls, weight_dir: str): """Load default values from model generation config.""" try: - generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() + generation_cfg = get_generation_config_diff_dict(weight_dir) cls._loaded_defaults = { "do_sample": generation_cfg.get("do_sample", True), "presence_penalty": generation_cfg.get("presence_penalty", 0.0), @@ -242,7 +243,7 @@ class ChatCompletionRequest(BaseModel): def load_generation_cfg(cls, weight_dir: str): """Load default values from model generation config.""" try: - generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() + generation_cfg = get_generation_config_diff_dict(weight_dir) cls._loaded_defaults = { "do_sample": generation_cfg.get("do_sample", True), "presence_penalty": generation_cfg.get("presence_penalty", 0.0), diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index cbc63c898d..3e1954502c 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -4,7 +4,7 @@ """ import os from typing import List, Optional, Union, Tuple -from transformers import GenerationConfig +from lightllm.utils.config_utils import get_generation_config_diff_dict from lightllm.server.req_id_generator import MAX_BEST_OF @@ -110,7 +110,7 @@ def __init__( @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() + generation_cfg = get_generation_config_diff_dict(weight_dir) cls._do_sample = generation_cfg.get("do_sample", False) cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0) cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index c39559f5f6..cea422d249 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -1,7 +1,7 @@ import os import ctypes from typing import Optional, List, Tuple, Union -from transformers import GenerationConfig +from lightllm.utils.config_utils import get_generation_config_diff_dict from lightllm.server.req_id_generator import MAX_BEST_OF from .pd_kv_trans_params import PDKVTransParamObj @@ -395,7 +395,7 @@ def init(self, tokenizer, **kwargs): @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() + generation_cfg = get_generation_config_diff_dict(weight_dir) cls._do_sample = generation_cfg.get("do_sample", False) cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0) cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0) diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index c892d4b0bc..3c6109829c 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -1,6 +1,6 @@ import json import os -from typing import Optional, List +from typing import Any, Dict, Optional, List from functools import lru_cache from .envs_utils import get_env_start_args from lightllm.utils.log_utils import init_logger @@ -14,6 +14,13 @@ def get_config_json(model_path: str): return json_obj +def get_generation_config_diff_dict(model_path: str) -> Dict[str, Any]: + from transformers import GenerationConfig + + generation_cfg = GenerationConfig.from_pretrained(model_path, trust_remote_code=True).to_diff_dict() + return {key: value for key, value in generation_cfg.items() if value is not None} + + def _derive_max_req_total_len_from_model_config(model_dir: str) -> Optional[int]: """ Derive `max_req_total_len` from model config.json. From bc6ad976b24d777d6d931c7c96a20140ee597569 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 3 Jul 2026 03:13:22 +0000 Subject: [PATCH 056/214] Optimize AI code --- lightllm/common/req_manager.py | 1 - .../server/router/model_infer/infer_batch.py | 194 +++++------------- .../model_infer/mode_backend/base_backend.py | 34 ++- 3 files changed, 72 insertions(+), 157 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 1af4991a3d..6023988dad 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -953,7 +953,6 @@ def concat_prompt_cache_payloads(self, payloads: List[DeepseekV4PromptCachePaylo def build_prompt_cache_payload( self, - req_idx: int, cache_len: int, ) -> DeepseekV4PromptCachePayload: """构造插入载荷。compressor 状态不进载荷(c4 随 swa 页生灭、c128 边界自然归零), diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 1d5db33bfb..7ce9d25056 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -124,7 +124,7 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: return req_objs def free_a_req_mem(self, free_token_index: List, req: "InferReq"): - is_dsv4_req_manager = hasattr(self.req_manager, "build_prompt_cache_payload") + is_dsv4_req_manager = isinstance(self.req_manager, DeepseekV4ReqManager) if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) if is_dsv4_req_manager: @@ -145,11 +145,6 @@ def free_a_req_mem(self, free_token_index: List, req: "InferReq"): req.shm_req.shm_cur_kv_len = req.cur_kv_len return - def _append_free_token_index(self, free_token_index: List, tensor: torch.Tensor): - if tensor.numel() > 0: - free_token_index.append(tensor) - return - def _full_att_free_req(self, free_token_index: List, req: "InferReq"): input_token_ids = req.get_input_token_ids() key = torch.tensor(input_token_ids[0 : req.cur_kv_len], dtype=torch.int64, device="cpu") @@ -166,59 +161,30 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): return def _dsv4_full_att_free_req(self, free_token_index: List, req: "InferReq"): - if req.cur_kv_len == 0: - free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0:0]) - return - old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len inserted_len = old_prefix_len duplicate_prefix_len = old_prefix_len - # 载荷只剩按页 bitmap(compressor 状态随 swa 页生灭/边界自然归零,不进载荷), - # 任意 128 对齐前缀皆可插入——含生成段(floor(cur_kv_len) 边界,回收保留尾页保证其驻留)。 cache_len = self.radix_cache.align_len(req.cur_kv_len) self.req_manager: DeepseekV4ReqManager if cache_len > old_prefix_len: - payload = self.req_manager.build_prompt_cache_payload(req.req_idx, cache_len) + payload = self.req_manager.build_prompt_cache_payload(cache_len) value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() - # 按页有效性 bitmap 用插入时刻的映射写定(此后只会被阀清 0,不会复活)。水位线 - # 纯 CPU 推导,避免 router 关键路径上的 GPU gather 同步(每插入一次要等全部在途 - # decode kernel)。插入门: 截掉结尾的 invalid 页 —— 它们生来不可命中,还会永久 - # 挡住后续更长前缀复用同一段 token(全量重插会因前缀已存在而保留旧 bitmap)。 - page_size = self.req_manager.get_prompt_cache_page_size() - bitmap = self.req_manager.swa_page_valid_from_watermark(req.req_idx, cache_len) - n_pages = int(bitmap.numel()) - while n_pages > 0 and not bool(bitmap[n_pages - 1]): - n_pages -= 1 - gated_len = n_pages * page_size - if gated_len < cache_len: - logger.info( - f"DeepSeek-V4 prompt cache insert gate: trailing swa pages already evicted, " - f"shrink insert {cache_len} -> {gated_len}" - ) - cache_len = gated_len - payload.cache_len = cache_len - payload.swa_page_valid = bitmap[:n_pages].clone() - - if cache_len > old_prefix_len: - input_token_ids = req.get_input_token_ids() - key = torch.tensor(input_token_ids[0:cache_len], dtype=torch.int64, device="cpu") - duplicate_prefix_len, cache_node = self.radix_cache.insert(key, value[:cache_len], extra_value=payload) - inserted_len = 0 if cache_node is None else cache_node.node_prefix_total_len - if inserted_len != cache_len: - inserted_len = old_prefix_len - duplicate_prefix_len = old_prefix_len + + payload.swa_page_valid = self.req_manager.swa_page_valid_from_watermark(req.req_idx, cache_len) + + key = torch.tensor(req.get_input_token_ids()[0:cache_len], dtype=torch.int64, device="cpu") + duplicate_prefix_len, _ = self.radix_cache.insert(key, value[:cache_len], extra_value=payload) + inserted_len = cache_len dense_row = self.req_manager.req_to_token_indexs[req.req_idx] - self._append_free_token_index(free_token_index, dense_row[old_prefix_len:duplicate_prefix_len]) - self._append_free_token_index(free_token_index, dense_row[inserted_len : req.cur_kv_len]) + if duplicate_prefix_len > old_prefix_len: + free_token_index.append(dense_row[old_prefix_len:duplicate_prefix_len]) + if req.cur_kv_len > inserted_len: + free_token_index.append(dense_row[inserted_len : req.cur_kv_len]) if len(free_token_index) == 0: free_token_index.append(dense_row[0:0]) - # 释放的 full 槽经 mem_manager.free 级联回收 swa/c4/c128(映射键控,无需收集槽位)。 - # pause 路径不会走 req_manager.free/init: 复位出窗水位线(残留水位线会破坏下一次 - # prefill 的共享前缀保护)并清 c128 在途状态(恢复命中走 extend 续算,若残留暂停前的 - # 半窗聚合会算错;c128 状态在 128 对齐命中边界本应为零)。 self.req_manager.init_compress_state(req.req_idx) if req.shared_kv_node is not None: @@ -438,8 +404,7 @@ def get_can_alloc_token_num(self): return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num def get_can_alloc_dsv4_swa_page_num(self): - mem_manager = self.req_manager.mem_manager - allocator = getattr(mem_manager, "swa_page_allocator", None) + allocator = getattr(self.req_manager.mem_manager, "swa_page_allocator", None) if allocator is None: return None @@ -448,44 +413,6 @@ def get_can_alloc_dsv4_swa_page_num(self): radix_cache_unref_page_num = self.radix_cache.get_unrefed_swa_pages_num() return int(allocator.can_use_mem_size) + radix_cache_unref_page_num - def get_dsv4_swa_prefill_need_page_num(self, req: "InferReq", is_chuncked_prefill: bool): - page_size = self._get_dsv4_swa_page_size() - if page_size is None: - return 0 - - start = int(req.cur_kv_len) - if is_chuncked_prefill: - end = int(req.get_chuncked_input_token_len()) - else: - end = int(req.get_cur_total_len()) - if end <= start: - return 0 - first_new_page = (start + page_size - 1) // page_size - last_page = (end - 1) // page_size - return last_page - first_new_page + 1 - - def get_dsv4_swa_decode_need_page_num(self, req: "InferReq"): - page_size = self._get_dsv4_swa_page_size() - if page_size is None: - return 0 - - seq_len = int(req.get_cur_total_len()) - if seq_len <= 0: - return 0 - return 1 if (seq_len - 1) % page_size == 0 else 0 - - def _get_dsv4_swa_page_size(self): - mem_manager = self.req_manager.mem_manager - allocator = getattr(mem_manager, "swa_page_allocator", None) - if allocator is None: - return None - return mem_manager.swa_pool.page_size - - # ---- DeepSeek-V4 compressed-pool (c4/c128) admission, mirror of the swa helpers above ---- - # c4 is paged for fp8_paged_mqa_logits: 64 c4-slots/page == 256 full tokens. The prompt-cache - # radix is 256-aligned (DSV4_PROMPT_CACHE_PAGE_SIZE) and c4 is NOT windowed, so reclaimable c4 - # pages derive exactly from the unref token count (`// 256`) — no separate counter needed - # (unlike swa, which is windowed). c128 is slot-based: 1 slot per 128 full tokens (`// 128`). def get_can_alloc_dsv4_c4_page_num(self): allocator = getattr(self.req_manager.mem_manager, "c4_page_allocator", None) if allocator is None: @@ -508,39 +435,6 @@ def get_can_alloc_dsv4_c128_slot_num(self): ) // 128 return int(allocator.can_use_mem_size) + int(radix_unref_slot_num) - def get_dsv4_c4_decode_need_page_num(self, req: "InferReq"): - if getattr(self.req_manager.mem_manager, "c4_page_allocator", None) is None: - return 0 - seq_len = int(req.get_cur_total_len()) - # 与 _scatter_c4_decode_slots 一致: 关组(seq%4==0)且组末 c4-entry 落页首(entry%64==0) -> 开新页 - if seq_len > 0 and seq_len % 4 == 0 and (seq_len // 4 - 1) % 64 == 0: - return 1 - return 0 - - def get_dsv4_c128_decode_need_slot_num(self, req: "InferReq"): - if getattr(self.req_manager.mem_manager, "c128_allocator", None) is None: - return 0 - seq_len = int(req.get_cur_total_len()) - return 1 if (seq_len > 0 and seq_len % 128 == 0) else 0 - - def get_dsv4_c4_prefill_need_page_num(self, req: "InferReq", is_chuncked_prefill: bool): - if getattr(self.req_manager.mem_manager, "c4_page_allocator", None) is None: - return 0 - start = int(req.cur_kv_len) - end = int(req.get_chuncked_input_token_len()) if is_chuncked_prefill else int(req.get_cur_total_len()) - first, last = start // 4, end // 4 - if last <= first: - return 0 - # 安全上界: 覆盖 c4-entry 区间 [first,last) 触及的全部 64-页(忽略已分配的延续页 -> 偏多, 安全) - return (last - 1) // 64 - first // 64 + 1 - - def get_dsv4_c128_prefill_need_slot_num(self, req: "InferReq", is_chuncked_prefill: bool): - if getattr(self.req_manager.mem_manager, "c128_allocator", None) is None: - return 0 - start = int(req.cur_kv_len) - end = int(req.get_chuncked_input_token_len()) if is_chuncked_prefill else int(req.get_cur_total_len()) - return max(0, end // 128 - start // 128) - def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """ 该函数用于在线性混合模型prefill后,如果存在大页匹配的情况下,将线性层状态复制到 @@ -753,6 +647,15 @@ def __init__( self.get_chuncked_input_token_len = self.get_chuncked_input_token_len_for_linear_att self.get_chuncked_input_token_ids = self.get_chuncked_input_token_ids_for_linear_att + mem_manager = g_infer_context.req_manager.mem_manager + self.dsv4_swa_page_size: Optional[int] = ( + mem_manager.swa_pool.page_size if getattr(mem_manager, "swa_pool", None) is not None else None + ) + self.dsv4_c4_page_size: Optional[int] = ( + mem_manager.c4_pool.page_size if getattr(mem_manager, "c4_pool", None) is not None else None + ) + self.dsv4_has_c128: bool = getattr(mem_manager, "c128_pool", None) is not None + self._init_all_state() self.generator = None @@ -992,7 +895,6 @@ def get_input_token_ids(self): def get_chuncked_input_token_ids(self): chunked_start = self.cur_kv_len chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) - chunked_end = self._align_chuncked_end_for_prompt_cache(chunked_start, chunked_end) return self.shm_req.shm_prompt_ids.arr[0:chunked_end] def get_chuncked_input_token_ids_for_linear_att(self): @@ -1013,23 +915,6 @@ def get_chuncked_input_token_ids_for_linear_att(self): def get_chuncked_input_token_len(self): chunked_start = self.cur_kv_len chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) - return self._align_chuncked_end_for_prompt_cache(chunked_start, chunked_end) - - def _align_chuncked_end_for_prompt_cache(self, chunked_start: int, chunked_end: int): - radix_cache = g_infer_context.radix_cache - page_size = getattr(radix_cache, "page_size", 1) if radix_cache is not None else 1 - if page_size <= 1 or self.sampling_param.disable_prompt_cache: - return chunked_end - prompt_end = int(self.shm_req.input_len) - chunked_start = int(chunked_start) - chunked_end = int(chunked_end) - if chunked_end >= prompt_end: - return chunked_end - - assert self.args.chunked_prefill_size % page_size == 0, ( - f"chunked_prefill_size={self.args.chunked_prefill_size} must be divisible by " - f"prompt-cache page_size={page_size}" - ) return chunked_end def get_chuncked_input_token_len_for_linear_att(self): @@ -1098,6 +983,41 @@ def _normal_decode_need_token_num(self) -> int: def _mtp_decode_need_token_num(self) -> int: return (1 + self.mtp_step) * 2 + def get_dsv4_prefill_need_page_and_slot_num(self, is_chuncked_prefill: bool) -> Tuple[int, int, int]: + start = self.cur_kv_len + end = self.get_chuncked_input_token_len() if is_chuncked_prefill else self.get_cur_total_len() + if end <= start: + return 0, 0, 0 + + swa_page_num = 0 + if self.dsv4_swa_page_size is not None: + first_new_page = (start + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size + last_page = (end - 1) // self.dsv4_swa_page_size + swa_page_num = last_page - first_new_page + 1 + + c4_page_num = 0 + if self.dsv4_c4_page_size is not None: + first, last = start // 4, end // 4 + if last > first: + # Safe upper bound: touched c4 pages, including a possible already-allocated continuation page. + c4_page_num = (last - 1) // self.dsv4_c4_page_size - first // self.dsv4_c4_page_size + 1 + + c128_slot_num = max(0, end // 128 - start // 128) if self.dsv4_has_c128 else 0 + return swa_page_num, c4_page_num, c128_slot_num + + def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: + seq_len = self.get_cur_total_len() + if seq_len <= 0: + return 0, 0, 0 + + swa_page_num = 1 if self.dsv4_swa_page_size is not None and (seq_len - 1) % self.dsv4_swa_page_size == 0 else 0 + c4_page_num = 0 + if self.dsv4_c4_page_size is not None and seq_len % 4 == 0: + entry = seq_len // 4 - 1 + c4_page_num = 1 if entry % self.dsv4_c4_page_size == 0 else 0 + c128_slot_num = 1 if self.dsv4_has_c128 and seq_len % 128 == 0 else 0 + return swa_page_num, c4_page_num, c128_slot_num + class InferReqUpdatePack: """ diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index 07eaeefa14..b555c4b80c 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -6,6 +6,7 @@ import torch.distributed as dist from typing import List, Tuple, Callable, Optional from transformers.configuration_utils import PretrainedConfig +from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel from lightllm.utils.infer_utils import set_random_seed from lightllm.utils.log_utils import init_logger from lightllm.models import get_model @@ -153,7 +154,7 @@ def init_model(self, kvargs): self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) - if hasattr(self.model.req_manager, "build_prompt_cache_payload"): + if isinstance(self.model, DeepseekV4TpPartModel): self.support_overlap = False if self.is_linear_att_mixed_model: @@ -166,7 +167,6 @@ def init_model(self, kvargs): if not self.use_dynamic_prompt_cache: self.radix_cache = None - setattr(self.args, "dynamic_prompt_cache_page_size", 1) else: if self.is_linear_att_mixed_model: self.radix_cache = LinearAttPagedRadixCache( @@ -178,14 +178,12 @@ def init_model(self, kvargs): kv_cache_mem_manager=self.model.mem_manager, linear_att_small_page_buffers=self.linear_att_cache_manager, ) - setattr(self.args, "dynamic_prompt_cache_page_size", 1) else: radix_page_size = 1 radix_extra_value_ops = None - if hasattr(self.model.req_manager, "get_prompt_cache_value_ops"): + if isinstance(self.model, DeepseekV4TpPartModel): radix_page_size = self.model.req_manager.get_prompt_cache_page_size() radix_extra_value_ops = self.model.req_manager.get_prompt_cache_value_ops() - setattr(self.args, "dynamic_prompt_cache_page_size", radix_page_size) self.radix_cache = RadixCache( unique_name=get_unique_server_name(), total_token_num=self.model.mem_manager.size, @@ -194,10 +192,15 @@ def init_model(self, kvargs): page_size=radix_page_size, extra_value_ops=radix_extra_value_ops, ) - if radix_extra_value_ops is not None and hasattr(self.model.mem_manager, "register_swa_free_hook"): - # swa 页 allocator 触底时让 radix 对 ref==0 节点 free swa 页(DeepSeek-V4)。 + if isinstance(self.model, DeepseekV4TpPartModel): self.model.mem_manager.register_swa_free_hook(self.radix_cache.free_unreferenced_swa_pages) + if not self.disable_chunked_prefill and radix_page_size > 1: + assert self.args.chunked_prefill_size % radix_page_size == 0, ( + f"chunked_prefill_size={self.args.chunked_prefill_size} must be divisible by " + f"prompt-cache page_size={radix_page_size}" + ) + if "prompt_cache_kv_buffer" in model_cfg: assert self.use_dynamic_prompt_cache self.preload_prompt_cache_kv_buffer(model_cfg) @@ -633,9 +636,7 @@ def _get_classed_reqs( if is_decode: token_num = req_obj.decode_need_token_num() - swa_page_num = g_infer_context.get_dsv4_swa_decode_need_page_num(req_obj) - c4_page_num = g_infer_context.get_dsv4_c4_decode_need_page_num(req_obj) - c128_slot_num = g_infer_context.get_dsv4_c128_decode_need_slot_num(req_obj) + swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_decode_need_page_and_slot_num() if ( token_num <= can_alloc_token_num and (can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num) @@ -661,17 +662,12 @@ def _get_classed_reqs( if req_obj.is_slave_req(): continue - token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=not self.disable_chunked_prefill) + is_chuncked_prefill = not self.disable_chunked_prefill + token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=is_chuncked_prefill) if prefill_tokens + token_num > self.batch_max_tokens: continue - swa_page_num = g_infer_context.get_dsv4_swa_prefill_need_page_num( - req_obj, is_chuncked_prefill=not self.disable_chunked_prefill - ) - c4_page_num = g_infer_context.get_dsv4_c4_prefill_need_page_num( - req_obj, is_chuncked_prefill=not self.disable_chunked_prefill - ) - c128_slot_num = g_infer_context.get_dsv4_c128_prefill_need_slot_num( - req_obj, is_chuncked_prefill=not self.disable_chunked_prefill + swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_prefill_need_page_and_slot_num( + is_chuncked_prefill=is_chuncked_prefill ) if ( token_num <= can_alloc_token_num From 99520f970e1449dded065b2eca4100463419ff56 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 3 Jul 2026 03:29:50 +0000 Subject: [PATCH 057/214] Optimize AI code --- .../server/router/model_infer/infer_batch.py | 88 ++++++++----------- .../model_infer/mode_backend/base_backend.py | 64 +++++++------- 2 files changed, 72 insertions(+), 80 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 7ce9d25056..ea12cc6fc0 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -40,6 +40,7 @@ class InferenceContext: overlap_stream: torch.cuda.Stream = None # 一些情况下推理进程进行异步折叠操作的异步流对象。 cpu_kv_cache_stream: torch.cuda.Stream = None # 用 cpu kv cache 操作的 stream is_linear_att_mixed_model: bool = False # 标记模型是否是full att 混合 linear att 的混合模型。 + is_deepseek_v4: bool = False def register( self, @@ -66,6 +67,7 @@ def register( self.vocab_size = vocab_size self.is_linear_att_mixed_model = isinstance(self.req_manager, ReqManagerForMamba) + self.is_deepseek_v4 = isinstance(self.req_manager, DeepseekV4ReqManager) return @@ -124,17 +126,16 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: return req_objs def free_a_req_mem(self, free_token_index: List, req: "InferReq"): - is_dsv4_req_manager = isinstance(self.req_manager, DeepseekV4ReqManager) if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) - if is_dsv4_req_manager: + if self.is_deepseek_v4: # 槽位随 full 槽经 mem_manager.free 级联回收。pause 路径不释放 req_idx, # 必须在此复位出窗水位线 + 清 c128 在途状态(恢复命中走 extend,不会再有 # restore/zero 时机;c4 状态随 swa 页生灭,无需处理)。 self.req_manager.init_compress_state(req.req_idx) else: if not self.is_linear_att_mixed_model: - if is_dsv4_req_manager: + if self.is_deepseek_v4: self._dsv4_full_att_free_req(free_token_index=free_token_index, req=req) else: self._full_att_free_req(free_token_index=free_token_index, req=req) @@ -403,37 +404,28 @@ def get_can_alloc_token_num(self): ) return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num - def get_can_alloc_dsv4_swa_page_num(self): - allocator = getattr(self.req_manager.mem_manager, "swa_page_allocator", None) - if allocator is None: - return None - + def get_can_alloc_dsv4_page_and_slot_num(self): + self.req_manager: DeepseekV4ReqManager + mem_manager = self.req_manager.mem_manager radix_cache_unref_page_num = 0 + radix_cache_unref_token_num = 0 if self.radix_cache is not None: radix_cache_unref_page_num = self.radix_cache.get_unrefed_swa_pages_num() - return int(allocator.can_use_mem_size) + radix_cache_unref_page_num - - def get_can_alloc_dsv4_c4_page_num(self): - allocator = getattr(self.req_manager.mem_manager, "c4_page_allocator", None) - if allocator is None: - return None - radix_unref_page_num = 0 - if self.radix_cache is not None: - radix_unref_page_num = ( - self.radix_cache.get_tree_total_tokens_num() - self.radix_cache.get_refed_tokens_num() - ) // 256 - return int(allocator.can_use_mem_size) + int(radix_unref_page_num) - - def get_can_alloc_dsv4_c128_slot_num(self): - allocator = getattr(self.req_manager.mem_manager, "c128_allocator", None) - if allocator is None: - return None - radix_unref_slot_num = 0 - if self.radix_cache is not None: - radix_unref_slot_num = ( + radix_cache_unref_token_num = ( self.radix_cache.get_tree_total_tokens_num() - self.radix_cache.get_refed_tokens_num() - ) // 128 - return int(allocator.can_use_mem_size) + int(radix_unref_slot_num) + ) + swa_page_num = int(mem_manager.swa_page_allocator.can_use_mem_size) + radix_cache_unref_page_num + + c4_page_num = 0 + if mem_manager.c4_page_allocator is not None: + c4_page_num = int(mem_manager.c4_page_allocator.can_use_mem_size) + int( + radix_cache_unref_token_num // self.req_manager.get_prompt_cache_page_size() + ) + + c128_slot_num = 0 + if mem_manager.c128_allocator is not None: + c128_slot_num = int(mem_manager.c128_allocator.can_use_mem_size) + int(radix_cache_unref_token_num // 128) + return swa_page_num, c4_page_num, c128_slot_num def copy_linear_att_state_to_cache_buffer(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """ @@ -647,14 +639,13 @@ def __init__( self.get_chuncked_input_token_len = self.get_chuncked_input_token_len_for_linear_att self.get_chuncked_input_token_ids = self.get_chuncked_input_token_ids_for_linear_att - mem_manager = g_infer_context.req_manager.mem_manager - self.dsv4_swa_page_size: Optional[int] = ( - mem_manager.swa_pool.page_size if getattr(mem_manager, "swa_pool", None) is not None else None - ) - self.dsv4_c4_page_size: Optional[int] = ( - mem_manager.c4_pool.page_size if getattr(mem_manager, "c4_pool", None) is not None else None - ) - self.dsv4_has_c128: bool = getattr(mem_manager, "c128_pool", None) is not None + if g_infer_context.is_deepseek_v4: + mem_manager = g_infer_context.req_manager.mem_manager + self.dsv4_swa_page_size: int = mem_manager.swa_pool.page_size + self.dsv4_c4_page_size: int = ( + mem_manager.c4_pool.page_size if mem_manager.c4_page_allocator is not None else 0 + ) + self.dsv4_has_c128: bool = mem_manager.c128_allocator is not None self._init_all_state() @@ -989,18 +980,15 @@ def get_dsv4_prefill_need_page_and_slot_num(self, is_chuncked_prefill: bool) -> if end <= start: return 0, 0, 0 - swa_page_num = 0 - if self.dsv4_swa_page_size is not None: - first_new_page = (start + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size - last_page = (end - 1) // self.dsv4_swa_page_size - swa_page_num = last_page - first_new_page + 1 + first_new_page = (start + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size + last_page = (end - 1) // self.dsv4_swa_page_size + swa_page_num = last_page - first_new_page + 1 c4_page_num = 0 - if self.dsv4_c4_page_size is not None: - first, last = start // 4, end // 4 - if last > first: - # Safe upper bound: touched c4 pages, including a possible already-allocated continuation page. - c4_page_num = (last - 1) // self.dsv4_c4_page_size - first // self.dsv4_c4_page_size + 1 + first, last = start // 4, end // 4 + if last > first: + # Safe upper bound: touched c4 pages, including a possible already-allocated continuation page. + c4_page_num = (last - 1) // self.dsv4_c4_page_size - first // self.dsv4_c4_page_size + 1 c128_slot_num = max(0, end // 128 - start // 128) if self.dsv4_has_c128 else 0 return swa_page_num, c4_page_num, c128_slot_num @@ -1010,9 +998,9 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: if seq_len <= 0: return 0, 0, 0 - swa_page_num = 1 if self.dsv4_swa_page_size is not None and (seq_len - 1) % self.dsv4_swa_page_size == 0 else 0 + swa_page_num = 1 if (seq_len - 1) % self.dsv4_swa_page_size == 0 else 0 c4_page_num = 0 - if self.dsv4_c4_page_size is not None and seq_len % 4 == 0: + if seq_len % 4 == 0: entry = seq_len // 4 - 1 c4_page_num = 1 if entry % self.dsv4_c4_page_size == 0 else 0 c128_slot_num = 1 if self.dsv4_has_c128 and seq_len % 128 == 0 else 0 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 b555c4b80c..5005a02f31 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -6,14 +6,13 @@ import torch.distributed as dist from typing import List, Tuple, Callable, Optional from transformers.configuration_utils import PretrainedConfig -from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel 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.server.router.model_infer.infer_batch import InferReq, InferReqUpdatePack from lightllm.server.router.token_load import TokenLoad from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.req_manager import ReqManagerForMamba +from lightllm.common.req_manager import DeepseekV4ReqManager, ReqManagerForMamba from lightllm.common.linear_att_cache_manager import LinearAttCacheManager from lightllm.server.router.dynamic_prompt.linear_att_radix_cache import LinearAttPagedRadixCache from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache @@ -154,7 +153,8 @@ def init_model(self, kvargs): self.model: TpPartBaseModel = self.model # for easy typing set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) - if isinstance(self.model, DeepseekV4TpPartModel): + self.is_deepseek_v4 = isinstance(self.model.req_manager, DeepseekV4ReqManager) + if self.is_deepseek_v4: self.support_overlap = False if self.is_linear_att_mixed_model: @@ -181,7 +181,7 @@ def init_model(self, kvargs): else: radix_page_size = 1 radix_extra_value_ops = None - if isinstance(self.model, DeepseekV4TpPartModel): + if self.is_deepseek_v4: radix_page_size = self.model.req_manager.get_prompt_cache_page_size() radix_extra_value_ops = self.model.req_manager.get_prompt_cache_value_ops() self.radix_cache = RadixCache( @@ -192,7 +192,7 @@ def init_model(self, kvargs): page_size=radix_page_size, extra_value_ops=radix_extra_value_ops, ) - if isinstance(self.model, DeepseekV4TpPartModel): + if self.is_deepseek_v4: self.model.mem_manager.register_swa_free_hook(self.radix_cache.free_unreferenced_swa_pages) if not self.disable_chunked_prefill and radix_page_size > 1: @@ -600,9 +600,13 @@ def _get_classed_reqs( prefill_tokens = 0 can_alloc_token_num = g_infer_context.get_can_alloc_token_num() - can_alloc_dsv4_swa_page_num = g_infer_context.get_can_alloc_dsv4_swa_page_num() - can_alloc_dsv4_c4_page_num = g_infer_context.get_can_alloc_dsv4_c4_page_num() - can_alloc_dsv4_c128_slot_num = g_infer_context.get_can_alloc_dsv4_c128_slot_num() + is_deepseek_v4 = self.is_deepseek_v4 + if is_deepseek_v4: + ( + can_alloc_dsv4_swa_page_num, + can_alloc_dsv4_c4_page_num, + can_alloc_dsv4_c128_slot_num, + ) = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() for req_obj in ready_reqs: @@ -636,20 +640,20 @@ def _get_classed_reqs( if is_decode: token_num = req_obj.decode_need_token_num() - swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_decode_need_page_and_slot_num() - if ( - token_num <= can_alloc_token_num - and (can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num) - and (can_alloc_dsv4_c4_page_num is None or c4_page_num <= can_alloc_dsv4_c4_page_num) - and (can_alloc_dsv4_c128_slot_num is None or c128_slot_num <= can_alloc_dsv4_c128_slot_num) - ): + can_run = token_num <= can_alloc_token_num + if can_run and is_deepseek_v4: + swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_decode_need_page_and_slot_num() + can_run = ( + swa_page_num <= can_alloc_dsv4_swa_page_num + and c4_page_num <= can_alloc_dsv4_c4_page_num + and c128_slot_num <= can_alloc_dsv4_c128_slot_num + ) + if can_run: decode_reqs.append(req_obj) can_alloc_token_num -= token_num - if can_alloc_dsv4_swa_page_num is not None: + if is_deepseek_v4: can_alloc_dsv4_swa_page_num -= swa_page_num - if can_alloc_dsv4_c4_page_num is not None: can_alloc_dsv4_c4_page_num -= c4_page_num - if can_alloc_dsv4_c128_slot_num is not None: can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: @@ -666,23 +670,23 @@ def _get_classed_reqs( token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=is_chuncked_prefill) if prefill_tokens + token_num > self.batch_max_tokens: continue - swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_prefill_need_page_and_slot_num( - is_chuncked_prefill=is_chuncked_prefill - ) - if ( - token_num <= can_alloc_token_num - and (can_alloc_dsv4_swa_page_num is None or swa_page_num <= can_alloc_dsv4_swa_page_num) - and (can_alloc_dsv4_c4_page_num is None or c4_page_num <= can_alloc_dsv4_c4_page_num) - and (can_alloc_dsv4_c128_slot_num is None or c128_slot_num <= can_alloc_dsv4_c128_slot_num) - ): + can_run = token_num <= can_alloc_token_num + if can_run and is_deepseek_v4: + swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_prefill_need_page_and_slot_num( + is_chuncked_prefill=is_chuncked_prefill + ) + can_run = ( + swa_page_num <= can_alloc_dsv4_swa_page_num + and c4_page_num <= can_alloc_dsv4_c4_page_num + and c128_slot_num <= can_alloc_dsv4_c128_slot_num + ) + if can_run: prefill_tokens += token_num prefill_reqs.append(req_obj) can_alloc_token_num -= token_num - if can_alloc_dsv4_swa_page_num is not None: + if is_deepseek_v4: can_alloc_dsv4_swa_page_num -= swa_page_num - if can_alloc_dsv4_c4_page_num is not None: can_alloc_dsv4_c4_page_num -= c4_page_num - if can_alloc_dsv4_c128_slot_num is not None: can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: From 57bf22dc27687522a8ae5697ceca7eabc53bf36d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 6 Jul 2026 05:46:53 +0000 Subject: [PATCH 058/214] convert list to text --- lightllm/models/deepseek_v4/model.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index d3f0715b0b..50e2e22de1 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -257,6 +257,11 @@ def apply_chat_template( # missing. Re-serialize dicts back to JSON strings so the encoder emits one # per real arg. for msg in msgs: + content = msg.get("content") + if isinstance(content, list) and all( + isinstance(part, dict) and part.get("type") == "text" for part in content + ): + msg["content"] = "".join(part.get("text") or "" for part in content) for tc in msg.get("tool_calls") or []: fn = tc.get("function") if isinstance(fn, dict) and isinstance(fn.get("arguments"), dict): From a61d9bcf72362d3d5d23d1375afdfad51d927bba Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 6 Jul 2026 06:07:11 +0000 Subject: [PATCH 059/214] support DS4 cache budgets in recover_paused_reqs --- .../server/router/model_infer/infer_batch.py | 30 +++++++++++++++++++ .../model_infer/mode_backend/base_backend.py | 7 ++++- 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 1d5db33bfb..24ecda7376 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -408,6 +408,9 @@ def recover_paused_reqs( paused_reqs: List["InferReq"], is_master_in_dp: bool, can_alloc_token_num: int, + can_alloc_dsv4_swa_page_num: int = None, + can_alloc_dsv4_c4_page_num: int = None, + can_alloc_dsv4_c128_slot_num: int = None, ): if paused_reqs: @@ -416,6 +419,27 @@ def recover_paused_reqs( if prefill_need_token_num > can_alloc_token_num: break + if can_alloc_dsv4_swa_page_num is not None: + swa_page_num = self.get_dsv4_swa_prefill_need_page_num(req, is_chuncked_prefill=False) + if swa_page_num > can_alloc_dsv4_swa_page_num: + break + else: + swa_page_num = 0 + + if can_alloc_dsv4_c4_page_num is not None: + c4_page_num = self.get_dsv4_c4_prefill_need_page_num(req, is_chuncked_prefill=False) + if c4_page_num > can_alloc_dsv4_c4_page_num: + break + else: + c4_page_num = 0 + + if can_alloc_dsv4_c128_slot_num is not None: + c128_slot_num = self.get_dsv4_c128_prefill_need_slot_num(req, is_chuncked_prefill=False) + if c128_slot_num > can_alloc_dsv4_c128_slot_num: + break + else: + c128_slot_num = 0 + if g_infer_context.is_linear_att_mixed_model: req._linear_match_radix_cache() else: @@ -427,6 +451,12 @@ def recover_paused_reqs( req.shm_req.is_paused = False logger.debug(f"infer recover paused req id {req.req_id}") can_alloc_token_num -= prefill_need_token_num + if can_alloc_dsv4_swa_page_num is not None: + can_alloc_dsv4_swa_page_num -= swa_page_num + if can_alloc_dsv4_c4_page_num is not None: + can_alloc_dsv4_c4_page_num -= c4_page_num + if can_alloc_dsv4_c128_slot_num is not None: + can_alloc_dsv4_c128_slot_num -= c128_slot_num return def get_can_alloc_token_num(self): 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 07eaeefa14..f85ec2cee3 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -707,7 +707,12 @@ def _get_classed_reqs( if recover_paused: g_infer_context.recover_paused_reqs( - paused_reqs=paused_reqs, is_master_in_dp=self.is_master_in_dp, can_alloc_token_num=can_alloc_token_num + paused_reqs=paused_reqs, + is_master_in_dp=self.is_master_in_dp, + can_alloc_token_num=can_alloc_token_num, + can_alloc_dsv4_swa_page_num=can_alloc_dsv4_swa_page_num, + can_alloc_dsv4_c4_page_num=can_alloc_dsv4_c4_page_num, + can_alloc_dsv4_c128_slot_num=can_alloc_dsv4_c128_slot_num, ) # 在 enable_prefill_decode_mixed 模式下,如果存在 prefill 请求和 decode 请求, From 7dd651abddf1579a3379d469876c0b2f692a79a0 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 6 Jul 2026 07:59:19 +0000 Subject: [PATCH 060/214] add error info --- lightllm/server/build_prompt.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index bda1aafea5..a28e51dd9e 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -169,5 +169,11 @@ async def build_prompt(request, tools) -> str: try: input_str = tokenizer.apply_chat_template(**kwargs, tokenize=False, add_generation_prompt=True, tools=tools) except Exception as e: - raise ValueError(f"Failed to build prompt: {e}") from None + logger.exception( + "Failed to build prompt. request=%s tools=%s template_kwargs=%s", + json.dumps(request.model_dump(by_alias=True, exclude_none=True), ensure_ascii=False, default=str), + json.dumps(tools, ensure_ascii=False, default=str), + json.dumps(kwargs, ensure_ascii=False, default=str), + ) + raise ValueError(f"Failed to build prompt: {e}") from e return input_str From 09402d2121727cfbef43219b37c2335415a62452 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 7 Jul 2026 01:44:36 +0000 Subject: [PATCH 061/214] support mtp --- lightllm/common/basemodel/basemodel.py | 60 ++++---- lightllm/common/basemodel/batch_objs.py | 39 +++++ .../deepseek4_mem_manager.py | 9 +- lightllm/common/req_manager.py | 51 +++++-- lightllm/models/deepseek_v4_mtp/__init__.py | 0 .../deepseek_v4_mtp/layer_infer/__init__.py | 0 .../layer_infer/pre_layer_infer.py | 54 +++++++ .../layer_infer/transformer_layer_infer.py | 8 + .../deepseek_v4_mtp/layer_weights/__init__.py | 0 .../pre_and_post_layer_weight.py | 107 +++++++++++++ .../layer_weights/transformer_layer_weight.py | 8 + lightllm/models/deepseek_v4_mtp/model.py | 143 ++++++++++++++++++ lightllm/server/api_start.py | 2 - .../model_infer/mode_backend/base_backend.py | 4 + .../mode_backend/chunked_prefill/impl.py | 19 ++- .../mode_backend/dp_backend/impl.py | 79 ++++++---- 16 files changed, 498 insertions(+), 85 deletions(-) create mode 100644 lightllm/models/deepseek_v4_mtp/__init__.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_infer/__init__.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_infer/pre_layer_infer.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_infer/transformer_layer_infer.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_weights/__init__.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/deepseek_v4_mtp/layer_weights/transformer_layer_weight.py create mode 100644 lightllm/models/deepseek_v4_mtp/model.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index f440b98213..be64fed55e 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -42,6 +42,8 @@ class TpPartBaseModel: + is_mtp_draft_model = False + # weight class pre_and_post_weight_class = None transformer_weight_class = None @@ -291,15 +293,6 @@ def _init_custom(self): @torch.no_grad() def forward(self, model_input: ModelInput): - # decode 槽位 prep: 放在 to_cuda 前, 优先使用 b_req_idx/b_seq_len 的 CPU mirror, - # 且此刻已在 forward 的 CUDA stream 上 -> 与后续 attention 同流, 无跨流竞态、无 D2H。 - # mem_indexes_cpu is None 时跳过: cudagraph warmup 的输入全在 CUDA 且 b_req_idx 全为 HOLD, prep 本就是 no-op。 - if not model_input.is_prefill and model_input.mem_indexes_cpu is not None: - self.req_manager.prepare_decode( - model_input.b_req_idx_cpu, - model_input.b_seq_len_cpu, - model_input.mem_indexes_cpu, - ) model_input.to_cuda() assert model_input.mem_indexes.is_cuda @@ -308,6 +301,20 @@ def forward(self, model_input: ModelInput): else: return self._decode(model_input) + def _prepare_decode_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor): + if model_input.mem_indexes_cpu is None: + return + if model_input.mtp_decode_slot_prepare_indices == (): + return + self.req_manager.prepare_decode( + model_input.b_req_idx_cpu, + model_input.b_seq_len_cpu, + model_input.b_mtp_index_cpu, + mem_indexes, + model_input.mtp_decode_slot_prepare_indices, + ) + return + def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() infer_state.input_ids = model_input.input_ids @@ -379,6 +386,9 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mtp_index = F.pad( new_model_input.b_mtp_index, (0, padded_batch_size), mode="constant", value=0 ) + new_model_input.b_mtp_index_cpu = F.pad( + new_model_input.b_mtp_index_cpu, (0, padded_batch_size), mode="constant", value=0 + ) new_model_input.b_seq_len = F.pad(new_model_input.b_seq_len, (0, padded_batch_size), mode="constant", value=2) new_model_input.b_req_idx_cpu = F.pad( new_model_input.b_req_idx_cpu, @@ -444,6 +454,7 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.b_req_idx, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID ) new_model_input.b_mtp_index = F.pad(new_model_input.b_mtp_index, (0, 1), mode="constant", value=0) + new_model_input.b_mtp_index_cpu = F.pad(new_model_input.b_mtp_index_cpu, (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_req_idx_cpu = F.pad( @@ -548,7 +559,7 @@ def _prefill( alloc_mem_index=infer_state.mem_index, max_q_seq_len=infer_state.max_q_seq_len, ) - if model_input.b_req_idx_cpu is not None: + if model_input.b_req_idx_cpu is not None and not self.is_mtp_draft_model: self.req_manager.prepare_prefill( b_req_idx=infer_state.b_req_idx, b_ready_cache_len=infer_state.b_ready_cache_len, @@ -584,6 +595,7 @@ def _decode( model_input.b_mtp_index, ) + origin_model_input = model_input origin_batch_size = model_input.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_ @@ -604,6 +616,7 @@ def _decode( infer_state.b_seq_len, infer_state.mem_index, ) + self._prepare_decode_slots(origin_model_input, infer_state.mem_index[:origin_batch_size]) infer_state.init_some_extra_state(self) infer_state.init_att_state() @@ -625,6 +638,7 @@ def _decode( infer_state.b_seq_len, infer_state.mem_index, ) + self._prepare_decode_slots(origin_model_input, infer_state.mem_index[:origin_batch_size]) infer_state.init_some_extra_state(self) infer_state.init_att_state() model_output = self._token_forward(infer_state) @@ -777,7 +791,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state0.mem_index, max_q_seq_len=infer_state0.max_q_seq_len, ) - if model_input0.b_req_idx_cpu is not None: + if model_input0.b_req_idx_cpu is not None and not self.is_mtp_draft_model: self.req_manager.prepare_prefill( b_req_idx=infer_state0.b_req_idx, b_ready_cache_len=infer_state0.b_ready_cache_len, @@ -799,7 +813,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state1.mem_index, max_q_seq_len=infer_state1.max_q_seq_len, ) - if model_input1.b_req_idx_cpu is not None: + if model_input1.b_req_idx_cpu is not None and not self.is_mtp_draft_model: self.req_manager.prepare_prefill( b_req_idx=infer_state1.b_req_idx, b_ready_cache_len=infer_state1.b_ready_cache_len, @@ -835,14 +849,6 @@ 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): - # decode 槽位 prep: 在 to_cuda 前使用 CPU mirror, 且已在 forward 的 CUDA stream 上 (见 forward 注释)。 - for mi in (model_input0, model_input1): - if mi.mem_indexes_cpu is not None: - self.req_manager.prepare_decode( - mi.b_req_idx_cpu, - mi.b_seq_len_cpu, - mi.mem_indexes_cpu, - ) model_input0.to_cuda() model_input1.to_cuda() assert self.args.enable_tpsp_mix_mode @@ -865,6 +871,8 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode assert model_input1.mem_indexes.is_cuda origin_batch_size = model_input0.batch_size + origin_model_input0 = model_input0 + origin_model_input1 = model_input1 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_ @@ -881,6 +889,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.b_seq_len, infer_state0.mem_index, ) + self._prepare_decode_slots(origin_model_input0, infer_state0.mem_index[:origin_batch_size]) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -891,6 +900,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.b_seq_len, infer_state1.mem_index, ) + self._prepare_decode_slots(origin_model_input1, infer_state1.mem_index[:origin_batch_size]) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -922,6 +932,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.b_seq_len, infer_state0.mem_index, ) + self._prepare_decode_slots(origin_model_input0, infer_state0.mem_index[:origin_batch_size]) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -932,6 +943,7 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.b_seq_len, infer_state1.mem_index, ) + self._prepare_decode_slots(origin_model_input1, infer_state1.mem_index[:origin_batch_size]) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -1233,13 +1245,7 @@ def _init_padded_req(self): def _gen_special_model_input(self, token_num: int): special_model_input = {} - is_mtp_draft_model = ( - "Deepseek3MTPModel" in str(self.__class__) - or "Qwen3MOEMTPModel" in str(self.__class__) - or "MistralMTPModel" in str(self.__class__) - or "Glm4MoeLiteMTPModel" in str(self.__class__) - ) - 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 7110d4895a..7a4d837081 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -1,3 +1,4 @@ +import copy import torch from dataclasses import dataclass, field from typing import Optional @@ -43,6 +44,7 @@ class ModelInput: # cpu 变量 mem_indexes_cpu: torch.Tensor = None b_req_idx_cpu: torch.Tensor = None + b_mtp_index_cpu: torch.Tensor = None b_seq_len_cpu: torch.Tensor = None b_ready_cache_len_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 @@ -55,6 +57,8 @@ class ModelInput: # mtp_draft_input_hiddens 用于模型 mtp 模式下 # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None + # 主模型为 None: 准备所有 MTP 列;draft 首轮为 (): 无新槽;draft 追加后为 (k,): 只准备新槽。 + mtp_decode_slot_prepare_indices: Optional[tuple] = None def _capture_cpu_mirror(self, tensor_name: str, mirror_name: str): tensor = getattr(self, tensor_name) @@ -64,10 +68,45 @@ def _capture_cpu_mirror(self, tensor_name: str, mirror_name: str): def capture_cpu_mirrors(self): self._capture_cpu_mirror("b_req_idx", "b_req_idx_cpu") + self._capture_cpu_mirror("b_mtp_index", "b_mtp_index_cpu") self._capture_cpu_mirror("b_seq_len", "b_seq_len_cpu") self._capture_cpu_mirror("b_ready_cache_len", "b_ready_cache_len_cpu") return + def make_mtp_draft_input(self): + model_input = copy.copy(self) + model_input.b_seq_len = self.b_seq_len.clone() + model_input.b_seq_len_cpu = self.b_seq_len_cpu.clone() + model_input.mtp_decode_slot_prepare_indices = () + return model_input + + def advance_mtp_decode_step( + self, + new_mem_indexes_cpu: torch.Tensor, + new_mem_indexes: torch.Tensor, + max_mtp_index: int, + ): + self.b_seq_len += 1 + self.b_seq_len_cpu += 1 + self.max_kv_seq_len += 1 + self.mtp_decode_slot_prepare_indices = (max_mtp_index,) + slots_per_req = max_mtp_index + 1 + self.mem_indexes_cpu = torch.cat( + [ + self.mem_indexes_cpu.view(-1, slots_per_req)[:, 1:], + new_mem_indexes_cpu.view(-1, 1), + ], + dim=1, + ).view(-1) + self.mem_indexes = torch.cat( + [ + self.mem_indexes.view(-1, slots_per_req)[:, 1:], + new_mem_indexes.view(-1, 1), + ], + dim=1, + ).view(-1) + return + def to_cuda(self): if self.input_ids is not None: self.input_ids = self.input_ids.cuda(non_blocking=True) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 784065d964..b37b18eee4 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -444,10 +444,10 @@ def alloc_swa_decode( req_to_token_indexs: torch.Tensor, ) -> None: """decode prep: 本步 token(位置 seq-1)的 swa 槽。整页起点开新页,否则上一 token 槽 +1 - (位置对齐不变式保证同页连续)。scatter 目标用 mem_indexes(此刻 req_to_token 尚未写本步)。 + (位置对齐不变式保证同页连续)。scatter 目标用当前步 mem_indexes。 - 注意: 续槽从上一位置的映射派生,故同一请求的多行(MTP 多 token/步)在同一批内不支持 - (DSV4 启动参数已拒绝 MTP;支持需按步内顺序分段派生)。""" + 注意: 续槽从上一位置的映射派生,故同一请求的多行(MTP 多 token/步)需要调用方按 + b_mtp_index 分段准备。""" page = DSV4_SWA_PAGE_SIZE hold_req_id = self.max_request_num req_list = b_req_idx_cpu.tolist() @@ -700,9 +700,6 @@ def pack_mla_kv_to_cache_fused_norm_rope( from lightllm.third_party.sglang_jit.dsv4 import fused_k_norm_rope_flashmla swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] - # 未映射槽位(-1, 如 decode 图 warmup 的 HOLD 行: prep 跳过 alloc_swa)对老 triton - # 写入核是显式 no-op;sglang fused 核无负槽位防护(负页偏移=非法访存),mask 到 - # swa HOLD 槽(垃圾桶语义,与 padding 行写入一致)。 swa_slots = torch.where(swa_slots < 0, torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), swa_slots) fused_k_norm_rope_flashmla( kv=kv, diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 6023988dad..45ceef9013 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -174,9 +174,16 @@ def prepare_prefill( prefill KV 槽位 prep 的模型 (DeepSeek-V4) override。""" return - def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): - """每个 decode step 在 to_cuda 之前调用的钩子 (优先使用 CPU mirror, 且已在 forward 的 - CUDA stream 上)。基类 no-op; 需要 per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" + def prepare_decode( + self, + b_req_idx_cpu, + b_seq_len_cpu, + b_mtp_index_cpu, + mem_indexes, + mtp_decode_slot_prepare_indices, + ): + """每个 decode step 在 attention metadata 构建前调用的钩子。基类 no-op; 需要 + per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" return @@ -491,11 +498,37 @@ def prepare_prefill_swa( ) return - def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu): - """decode 每步槽位 prep: 先 swa 再 compress。由 BaseModel.forward / microbatch_overlap_decode - 在 to_cuda 之前调用 (CPU mirror + forward 的 CUDA stream); 不再放在 _decode 里。""" - self.prepare_decode_swa(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) - self.prepare_decode_compress_slots(b_req_idx_cpu, b_seq_len_cpu, mem_indexes_cpu) + def prepare_decode( + self, + b_req_idx_cpu, + b_seq_len_cpu, + b_mtp_index_cpu, + mem_indexes, + mtp_decode_slot_prepare_indices, + ): + """decode 每步槽位 prep: 先 swa 再 compress。由 BaseModel 在 copy_kv_index_to_req + 之后、attention metadata 构建前调用。""" + max_mtp_index = int(b_mtp_index_cpu.max().item()) + if mtp_decode_slot_prepare_indices is None: + steps = range(max_mtp_index + 1) + else: + steps = mtp_decode_slot_prepare_indices + + batch_size = b_mtp_index_cpu.shape[0] + slots_per_req = max_mtp_index + 1 + assert batch_size % slots_per_req == 0 + for step in steps: + rows = slice(step, batch_size, slots_per_req) + self.prepare_decode_swa( + b_req_idx_cpu[rows], + b_seq_len_cpu[rows], + mem_indexes[rows], + ) + self.prepare_decode_compress_slots( + b_req_idx_cpu[rows], + b_seq_len_cpu[rows], + mem_indexes[rows], + ) return def prepare_prefill( @@ -869,7 +902,7 @@ def prepare_decode_compress_slots( mem_indexes: torch.Tensor, ) -> None: """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 - 组末 full 槽即本步的 mem_index(此刻 req_to_token_indexs 尚未写入本步槽位)。 + 组末 full 槽即本步的 mem_index。 从 CPU 镜像读 seq_len/req_idx(host 算术,无 D2H);非关组步 rows 为空 => 不调 _scatter,零同步。""" if self.n_c4 == 0 and self.n_c128 == 0: return diff --git a/lightllm/models/deepseek_v4_mtp/__init__.py b/lightllm/models/deepseek_v4_mtp/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4_mtp/layer_infer/__init__.py b/lightllm/models/deepseek_v4_mtp/layer_infer/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4_mtp/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4_mtp/layer_infer/pre_layer_infer.py new file mode 100644 index 0000000000..66b53ff4aa --- /dev/null +++ b/lightllm/models/deepseek_v4_mtp/layer_infer/pre_layer_infer.py @@ -0,0 +1,54 @@ +import torch + +from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.deepseek_v4_mtp.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4MTPPreAndPostLayerWeight, +) +from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer + + +class DeepseekV4MTPPreLayerInfer(LlamaPreLayerInfer): + def __init__(self, network_config): + super().__init__(network_config) + self.eps_ = network_config["rms_norm_eps"] + self.hidden_size = network_config["hidden_size"] + self.hc_mult = network_config["hc_mult"] + return + + def _mtp_forward( + self, + input_embdings: torch.Tensor, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4MTPPreAndPostLayerWeight, + ): + input_embdings = input_embdings.masked_fill(infer_state.position_ids.eq(0).view(-1, 1), 0) + target = infer_state.mtp_draft_input_hiddens + + layer_weight.enorm_weight_(input=input_embdings, eps=self.eps_, out=input_embdings) + e_proj = layer_weight.e_proj_weight_.mm(input_embdings) + + target = target.view(-1, self.hc_mult, self.hidden_size).contiguous() + target = layer_weight.hnorm_weight_(input=target, eps=self.eps_, alloc_func=self.alloc_tensor) + h_proj = layer_weight.h_proj_weight_.mm(target.reshape(-1, self.hidden_size)) + h_proj = h_proj.view(-1, self.hc_mult, self.hidden_size) + + output = h_proj + e_proj.unsqueeze(1) + return output.reshape(output.shape[0], self.hc_mult * self.hidden_size) + + def context_forward( + self, + input_ids, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4MTPPreAndPostLayerWeight, + ): + input_embdings = super().context_forward(input_ids, infer_state, layer_weight) + return self._mtp_forward(input_embdings, infer_state, layer_weight) + + def token_forward( + self, + input_ids, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4MTPPreAndPostLayerWeight, + ): + input_embdings = super().token_forward(input_ids, infer_state, layer_weight) + return self._mtp_forward(input_embdings, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4_mtp/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4_mtp/layer_infer/transformer_layer_infer.py new file mode 100644 index 0000000000..651ef9fc24 --- /dev/null +++ b/lightllm/models/deepseek_v4_mtp/layer_infer/transformer_layer_infer.py @@ -0,0 +1,8 @@ +from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer + + +class DeepseekV4MTPTransformerLayerInfer(DeepseekV4TransformerLayerInfer): + def __init__(self, layer_num, network_config): + super().__init__(layer_num, network_config) + self.is_last_layer = True + return diff --git a/lightllm/models/deepseek_v4_mtp/layer_weights/__init__.py b/lightllm/models/deepseek_v4_mtp/layer_weights/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/lightllm/models/deepseek_v4_mtp/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/deepseek_v4_mtp/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..2b139d016b --- /dev/null +++ b/lightllm/models/deepseek_v4_mtp/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,107 @@ +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 +from lightllm.models.deepseek_v4.triton_kernel.quant_convert import dequant_fp8_block_to_bf16 + + +class DeepseekV4MTPPreAndPostLayerWeight(PreAndPostLayerWeight): + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config) + self.quant_cfg: Quantcfg = quant_cfg + + hidden = network_config["hidden_size"] + vocab = network_config["vocab_size"] + hc_mult = network_config["hc_mult"] + layer_idx = network_config["n_layer"] + prefix = "mtp.0" + + self.wte_weight_ = EmbeddingWeight( + dim=hidden, + vocab_size=vocab, + weight_name=f"{prefix}.emb.tok_emb.weight", + data_type=self.data_type_, + ) + self.lm_head_weight_ = LMHeadWeight( + dim=hidden, + vocab_size=vocab, + weight_name=f"{prefix}.head.weight", + data_type=self.data_type_, + ) + self.final_norm_weight_ = RMSNormWeight( + dim=hidden, + weight_name=f"{prefix}.norm.weight", + data_type=self.data_type_, + ) + + self.e_proj_weight_ = ROWMMWeight( + in_dim=hidden, + out_dims=[hidden], + weight_names=f"{prefix}.e_proj.weight", + data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(layer_idx, "e_proj"), + tp_rank=0, + tp_world_size=1, + ) + self.h_proj_weight_ = ROWMMWeight( + in_dim=hidden, + out_dims=[hidden], + weight_names=f"{prefix}.h_proj.weight", + data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(layer_idx, "h_proj"), + tp_rank=0, + tp_world_size=1, + ) + self.enorm_weight_ = RMSNormWeight( + dim=hidden, + weight_name=f"{prefix}.enorm.weight", + data_type=self.data_type_, + ) + self.hnorm_weight_ = RMSNormWeight( + dim=hidden, + weight_name=f"{prefix}.hnorm.weight", + data_type=self.data_type_, + ) + + self.hc_head_fn_ = ParameterWeight( + weight_name=f"{prefix}.hc_head_fn", + data_type=torch.float32, + weight_shape=(hc_mult, hc_mult * hidden), + ) + self.hc_head_base_ = ParameterWeight( + weight_name=f"{prefix}.hc_head_base", + data_type=torch.float32, + weight_shape=(hc_mult,), + ) + self.hc_head_scale_ = ParameterWeight( + weight_name=f"{prefix}.hc_head_scale", + data_type=torch.float32, + weight_shape=(1,), + ) + return + + def load_hf_weights(self, weights): + self._dequant_in_place(weights) + return super().load_hf_weights(weights) + + def _dequant_in_place(self, weights): + for attr in (self.e_proj_weight_, self.h_proj_weight_): + for weight_name, scale_name in zip(attr.weight_names, attr.weight_scale_names): + scale_key = weight_name[: -len(".weight")] + ".scale" + if scale_key not in weights: + continue + if scale_name is None: + weights[weight_name] = dequant_fp8_block_to_bf16(weights[weight_name], weights[scale_key]).to( + self.data_type_ + ) + else: + weights[scale_name] = weights[scale_key].to(torch.float32) + del weights[scale_key] + return diff --git a/lightllm/models/deepseek_v4_mtp/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4_mtp/layer_weights/transformer_layer_weight.py new file mode 100644 index 0000000000..3281bcb98c --- /dev/null +++ b/lightllm/models/deepseek_v4_mtp/layer_weights/transformer_layer_weight.py @@ -0,0 +1,8 @@ +from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight + + +class DeepseekV4MTPTransformerLayerWeight(DeepseekV4TransformerLayerWeight): + def _parse_config(self): + super()._parse_config() + self.prefix = "mtp.0" + return diff --git a/lightllm/models/deepseek_v4_mtp/model.py b/lightllm/models/deepseek_v4_mtp/model.py new file mode 100644 index 0000000000..ae91a1576a --- /dev/null +++ b/lightllm/models/deepseek_v4_mtp/model.py @@ -0,0 +1,143 @@ +import gc +import os +from typing import List + +import torch +from safetensors import safe_open +from tqdm import tqdm + +from lightllm.common.basemodel import TpPartBaseModel +from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel +from lightllm.models.deepseek_v4_mtp.layer_infer.pre_layer_infer import DeepseekV4MTPPreLayerInfer +from lightllm.models.deepseek_v4_mtp.layer_infer.transformer_layer_infer import ( + DeepseekV4MTPTransformerLayerInfer, +) +from lightllm.models.deepseek_v4_mtp.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4MTPPreAndPostLayerWeight, +) +from lightllm.models.deepseek_v4_mtp.layer_weights.transformer_layer_weight import ( + DeepseekV4MTPTransformerLayerWeight, +) +import lightllm.utils.petrel_helper as utils +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + + +class DeepseekV4MTPModel(DeepseekV4TpPartModel): + is_mtp_draft_model = True + + pre_and_post_weight_class = DeepseekV4MTPPreAndPostLayerWeight + pre_layer_infer_class = DeepseekV4MTPPreLayerInfer + transformer_weight_class = DeepseekV4MTPTransformerLayerWeight + transformer_layer_infer_class = DeepseekV4MTPTransformerLayerInfer + + def __init__(self, kvargs: dict): + self._pre_init(kvargs) + super().__init__(kvargs) + return + + def _pre_init(self, kvargs: dict): + self.main_model: TpPartBaseModel = kvargs.pop("main_model") + self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") + return + + def _init_custom(self): + self._freqs_cis_sliding = self.main_model._freqs_cis_sliding + self._freqs_cis_compress = self.main_model._freqs_cis_compress + self._cos_cached_sliding = self.main_model._cos_cached_sliding + self._sin_cached_sliding = self.main_model._sin_cached_sliding + self._cos_cached_compress = self.main_model._cos_cached_compress + self._sin_cached_compress = self.main_model._sin_cached_compress + self.dsv4_workspace = self.main_model.dsv4_workspace + for layer in self.layers_infer: + layer.freqs_cis = self._freqs_cis_compress if layer.compress_ratio else self._freqs_cis_sliding + layer.cos_compress_table = self._cos_cached_compress + layer.sin_compress_table = self._sin_cached_compress + return + + def _init_req_manager(self): + self.req_manager = self.main_model.req_manager + return + + def _init_mem_manager(self): + self.mem_manager = self.main_model.mem_manager + return + + def _init_weights(self, start_layer_index=None): + assert start_layer_index is None + mtp_layer_index = self.config["n_layer"] + 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( + mtp_layer_index, + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + ] + return + + def _init_infer_layer(self, start_layer_index=None): + assert start_layer_index is None + self.pre_infer = self.pre_layer_infer_class(network_config=self.config) + self.post_infer = self.post_layer_infer_class(network_config=self.config) + 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] + ) + self.layers_infer = [self.transformer_layer_infer_class(total_pre_layers_num, network_config=self.config)] + return + + def _init_some_value(self): + super()._init_some_value() + self.layers_num = 1 + return + + def _gen_special_model_input(self, token_num: int): + return { + "mtp_draft_input_hiddens": torch.randn( + token_num, + self.config["hc_mult"] * self.config["hidden_size"], + dtype=self.data_type, + device="cuda", + ) + } + + def _load_hf_weights(self): + index_file = os.path.join(self.weight_dir_, "model.safetensors.index.json") + assert utils.PetrelHelper.exists(index_file), "DeepSeek-V4 MTP requires model.safetensors.index.json." + weight_map = utils.PetrelHelper.load_json(index_file)["weight_map"] + mtp_keys_by_file = {} + for key, file_ in weight_map.items(): + if key.startswith("mtp.0."): + mtp_keys_by_file.setdefault(file_, []).append(key) + + candidate_files = sorted(mtp_keys_by_file.keys()) + assert len(candidate_files) > 0, "DeepSeek-V4 MTP weights with prefix mtp.0. were not found." + + loaded_key_count = 0 + desc = f"pid {os.getpid()} Loading DeepSeek-V4 MTP weights" + for file_ in tqdm(candidate_files, total=len(candidate_files), desc=desc): + weights = {} + with safe_open(os.path.join(self.weight_dir_, file_), "pt", "cpu") as f: + for key in mtp_keys_by_file[file_]: + weights[key] = f.get_tensor(key) + + loaded_key_count += len(weights) + self.pre_post_weight.load_hf_weights(weights) + for layer in self.trans_layers_weight: + layer.load_hf_weights(weights) + del weights + gc.collect() + + self.pre_post_weight.verify_load() + [weight.verify_load() for weight in self.trans_layers_weight] + logger.info(f"loaded DeepSeek-V4 MTP weights: {loaded_key_count} tensors") + return + + def autotune_layers(self): + return 1 diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index c13f562af9..a5d622ff74 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -120,8 +120,6 @@ def normal_or_p_d_start(args): raise NotImplementedError("DeepSeek-V4 currently supports only run_mode=normal in LightLLM.") if args.enable_cpu_cache or args.enable_disk_cache: raise NotImplementedError("DeepSeek-V4 CPU/disk KV cache is not supported yet.") - if args.mtp_mode is not None or args.mtp_draft_model_dir is not None or args.mtp_step != 0: - raise NotImplementedError("DeepSeek-V4 MTP/speculative decoding is not supported yet.") if args.enable_ep_moe: raise NotImplementedError("DeepSeek-V4 EP MoE is not supported yet; use TP for now.") if "prompt_cache_kv_buffer" in get_config_json(args.model_dir): 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 5005a02f31..c18b3a45a5 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -42,6 +42,7 @@ 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.deepseek_v4_mtp.model import DeepseekV4MTPModel 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 @@ -351,6 +352,9 @@ def init_mtp_draft_model(self, main_kvargs: dict): 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 == "deepseek_v4": + assert self.args.mtp_mode == "eagle_with_att" + self.draft_models.append(DeepseekV4MTPModel(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)) 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 40d1c175aa..512bf712fa 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 @@ -241,7 +241,7 @@ def decode_mtp( model_input, run_reqs = prepare_decode_inputs(decode_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): - b_mtp_index_cpu = model_input.b_mtp_index + b_mtp_index_cpu = model_input.b_mtp_index_cpu model_output = self.model.forward(model_input) next_token_ids, next_token_logprobs = sample(model_output.logits, run_reqs, self.eos_id) # verify the next_token_ids @@ -345,7 +345,7 @@ def _draft_decode_vanilla( b_req_mtp_start_loc: torch.Tensor, ): # share some inference info with the main model - draft_model_input = main_model_input + draft_model_input = main_model_input.make_mtp_draft_input() draft_model_output = main_model_output draft_next_token_ids = next_token_ids all_next_token_ids = [] @@ -387,7 +387,7 @@ def _draft_decode_eagle( 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_input = main_model_input.make_mtp_draft_input() draft_model_output = main_model_output draft_next_token_ids = next_token_ids all_next_token_ids = [] @@ -401,14 +401,13 @@ def _draft_decode_eagle( 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.b_seq_len_cpu += 1 - draft_model_input.max_kv_seq_len += 1 + eagle_mem_indexes_cpu_i = eagle_mem_indexes_cpu[_step * num_reqs : (_step + 1) * num_reqs] 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) + draft_model_input.advance_mtp_decode_step( + eagle_mem_indexes_cpu_i, + eagle_mem_indexes_i, + self.mtp_step, + ) 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] 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 b051a1a3ab..1860fdd5b9 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 @@ -433,7 +433,7 @@ def prefill_mtp(self, event_pack: OverlapEventPack, prefill_reqs: List[InferReq] def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): model_input, run_reqs, _ = padded_prepare_decode_inputs(decode_reqs) - b_mtp_index_cpu = model_input.b_mtp_index + b_mtp_index_cpu = model_input.b_mtp_index_cpu req_num = len(run_reqs) with torch.cuda.stream(g_infer_context.get_overlap_stream()): @@ -537,14 +537,13 @@ def _draft_decode_vanilla( ): all_next_token_ids = [] # share some inference info with the main model - draft_model_input = model_input + draft_model_input = model_input.make_mtp_draft_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) - all_next_token_ids.append(draft_next_token_ids_gpu) - # process the draft model output for draft_model_idx in range(self.mtp_step): @@ -578,7 +577,7 @@ def _draft_decode_eagle( ): all_next_token_ids = [] # share some inference info with the main model - draft_model_input = model_input + draft_model_input = model_input.make_mtp_draft_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") @@ -587,7 +586,6 @@ def _draft_decode_eagle( real_req_num = req_num // (self.mtp_step + 1) padded_req_num = model_input.batch_size // (self.mtp_step + 1) - real_req_num - eagle_mem_indexes_cpu = None 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) @@ -602,20 +600,25 @@ def _draft_decode_eagle( 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.b_seq_len_cpu += 1 - draft_model_input.max_kv_seq_len += 1 + eagle_mem_indexes_cpu_i = eagle_mem_indexes_cpu[_step * real_req_num : (_step + 1) * real_req_num] eagle_mem_indexes_i = eagle_mem_indexes[_step * real_req_num : (_step + 1) * real_req_num] + eagle_mem_indexes_cpu_i = F.pad( + input=eagle_mem_indexes_cpu_i, + pad=(0, padded_req_num), + mode="constant", + value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + ) 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_model_input.advance_mtp_decode_step( + eagle_mem_indexes_cpu_i, + eagle_mem_indexes_i, + self.mtp_step, + ) draft_next_token_ids_gpu = self._gen_argmax_token_ids(draft_model_output) all_next_token_ids.append(draft_next_token_ids_gpu) @@ -740,8 +743,8 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf ) = padded_overlap_prepare_decode_inputs(decode_reqs) req_num0, req_num1 = len(run_reqs0), len(run_reqs1) all_next_token_ids = [] - b_mtp_index_cpu0 = model_input0.b_mtp_index - b_mtp_index_cpu1 = model_input1.b_mtp_index + b_mtp_index_cpu0 = model_input0.b_mtp_index_cpu + b_mtp_index_cpu1 = model_input1.b_mtp_index_cpu with torch.cuda.stream(g_infer_context.get_overlap_stream()): model_output0, model_output1 = self.model.microbatch_overlap_decode(model_input0, model_input1) @@ -873,7 +876,8 @@ def _draft_decode_vanilla_overlap( 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_input0 = model_input0.make_mtp_draft_input() + draft_model_input1 = model_input1.make_mtp_draft_input() 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") @@ -931,7 +935,8 @@ def _draft_decode_eagle_overlap( 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_input0 = model_input0.make_mtp_draft_input() + draft_model_input1 = model_input1.make_mtp_draft_input() 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") @@ -951,6 +956,8 @@ def _draft_decode_eagle_overlap( 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_indexes_cpu0 = eagle_mem_indexes_cpu[0 : real_req_num0 * self.mtp_step] + eagle_mem_indexes_cpu1 = eagle_mem_indexes_cpu[real_req_num0 * self.mtp_step : real_req_num * self.mtp_step] 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] @@ -967,35 +974,45 @@ def _draft_decode_eagle_overlap( draft_model_input0, draft_model_input1 ) - draft_model_input0.b_seq_len += 1 - draft_model_input0.b_seq_len_cpu += 1 - draft_model_input0.max_kv_seq_len += 1 + eagle_mem_indexes_cpu_i = eagle_mem_indexes_cpu0[_step * real_req_num0 : (_step + 1) * real_req_num0] eagle_mem_indexes_i = eagle_mem_indexes0[_step * real_req_num0 : (_step + 1) * real_req_num0] + eagle_mem_indexes_cpu_i = F.pad( + input=eagle_mem_indexes_cpu_i, + pad=(0, padded_req_num0), + mode="constant", + value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + ) 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_input0.advance_mtp_decode_step( + eagle_mem_indexes_cpu_i, + eagle_mem_indexes_i, + self.mtp_step, + ) - draft_model_input1.b_seq_len += 1 - draft_model_input1.b_seq_len_cpu += 1 - draft_model_input1.max_kv_seq_len += 1 + eagle_mem_indexes_cpu_i = eagle_mem_indexes_cpu1[_step * real_req_num1 : (_step + 1) * real_req_num1] eagle_mem_indexes_i = eagle_mem_indexes1[_step * real_req_num1 : (_step + 1) * real_req_num1] + eagle_mem_indexes_cpu_i = F.pad( + input=eagle_mem_indexes_cpu_i, + pad=(0, padded_req_num1), + mode="constant", + value=g_infer_context.req_manager.mem_manager.HOLD_TOKEN_MEMINDEX, + ) 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_model_input1.advance_mtp_decode_step( + eagle_mem_indexes_cpu_i, + eagle_mem_indexes_i, + self.mtp_step, + ) draft_next_token_ids_gpu0 = self._gen_argmax_token_ids(draft_model_output0) draft_next_token_ids_gpu1 = self._gen_argmax_token_ids(draft_model_output1) From 1b7e95c92e8c62f5ab5ef7c7c23c3ebc16b6ec6c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 7 Jul 2026 06:16:23 +0000 Subject: [PATCH 062/214] fix --- lightllm/common/basemodel/basemodel.py | 73 ++++++++++++++----- .../layer_infer/post_layer_infer.py | 3 +- lightllm/models/deepseek_v4_mtp/model.py | 14 ++-- 3 files changed, 63 insertions(+), 27 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index be64fed55e..2f4a814e0c 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -695,14 +695,20 @@ 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) + mtp_main_output_hiddens = None + if isinstance(predict_logits, tuple): + predict_logits, mtp_main_output_hiddens = predict_logits model_output = ModelOutput(logits=predict_logits) # 特殊模型特殊模式的额外输出 if self.is_mtp_mode: - input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - if infer_state.need_dp_prefill_balance: - input_embs = infer_state._all_to_all_unbalance_get(data=input_embs) - model_output.mtp_main_output_hiddens = input_embs.contiguous() + if mtp_main_output_hiddens is not None: + model_output.mtp_main_output_hiddens = mtp_main_output_hiddens.contiguous() + else: + 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() # 在开启使用deepep的时候,需要调用clear_deepep_buffer做资源清理,没有启用的时候 # 该调用没有实际意义 @@ -721,16 +727,22 @@ def _token_forward(self, infer_state: InferStateInfo): input_embs: torch.Tensor = layer.token_forward(input_embs, infer_state, self.trans_layers_weight[i]) last_input_embs = self.post_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) - predict_logits: torch.Tensor = self.post_infer.token_forward( + predict_logits = self.post_infer.token_forward( last_input_embs, infer_state=infer_state, layer_weight=self.pre_post_weight ) + mtp_main_output_hiddens = None + if isinstance(predict_logits, tuple): + predict_logits, mtp_main_output_hiddens = predict_logits 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() + if mtp_main_output_hiddens is not None: + model_output.mtp_main_output_hiddens = mtp_main_output_hiddens.contiguous() + else: + input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) + model_output.mtp_main_output_hiddens = input_embs.contiguous() # 在 cuda graph 模式下,输出需要转为 no ref tensor, 加强mem pool 的复用,降低显存的使用。 if infer_state.is_cuda_graph: @@ -991,18 +1003,31 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) g_cache_manager.cache_env_out() + mtp_main_output_hiddens = None + mtp_main_output_hiddens1 = None + if isinstance(predict_logits, tuple): + predict_logits, mtp_main_output_hiddens = predict_logits + if isinstance(predict_logits1, tuple): + predict_logits1, mtp_main_output_hiddens1 = predict_logits1 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) - 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() + if mtp_main_output_hiddens is not None: + model_output.mtp_main_output_hiddens = mtp_main_output_hiddens.contiguous() + else: + input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) + if infer_state.need_dp_prefill_balance: + input_embs = infer_state._all_to_all_unbalance_get(data=input_embs) + model_output.mtp_main_output_hiddens = input_embs.contiguous() + if mtp_main_output_hiddens1 is not None: + model_output1.mtp_main_output_hiddens = mtp_main_output_hiddens1.contiguous() + else: + input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) + if infer_state.need_dp_prefill_balance: + input_embs1 = infer_state1._all_to_all_unbalance_get(data=input_embs1) + model_output1.mtp_main_output_hiddens = input_embs1.contiguous() return model_output, model_output1 @@ -1029,15 +1054,27 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: predict_logits, predict_logits1 = self.post_infer.overlap_tpsp_token_forward( last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) + mtp_main_output_hiddens = None + mtp_main_output_hiddens1 = None + if isinstance(predict_logits, tuple): + predict_logits, mtp_main_output_hiddens = predict_logits + if isinstance(predict_logits1, tuple): + predict_logits1, mtp_main_output_hiddens1 = predict_logits1 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() + if mtp_main_output_hiddens is not None: + model_output.mtp_main_output_hiddens = mtp_main_output_hiddens.contiguous() + else: + input_embs = self.pre_infer._tpsp_allgather(input=input_embs, infer_state=infer_state) + model_output.mtp_main_output_hiddens = input_embs.contiguous() + if mtp_main_output_hiddens1 is not None: + model_output1.mtp_main_output_hiddens = mtp_main_output_hiddens1.contiguous() + else: + input_embs1 = self.pre_infer._tpsp_allgather(input=input_embs1, infer_state=infer_state1) + model_output1.mtp_main_output_hiddens = input_embs1.contiguous() if infer_state.is_cuda_graph: model_output.to_no_ref_tensor() diff --git a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py index 8eddfb3b9d..0eb5cfc2b6 100644 --- a/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py @@ -24,4 +24,5 @@ def token_forward(self, input_embdings, infer_state: DeepseekV4InferStateInfo, l cfg.get("hc_eps", 1e-6), self.alloc_tensor, ) - return super().token_forward(collapsed, infer_state, layer_weight) + logits = super().token_forward(collapsed, infer_state, layer_weight) + return logits, input_embdings diff --git a/lightllm/models/deepseek_v4_mtp/model.py b/lightllm/models/deepseek_v4_mtp/model.py index ae91a1576a..829ece2cd7 100644 --- a/lightllm/models/deepseek_v4_mtp/model.py +++ b/lightllm/models/deepseek_v4_mtp/model.py @@ -71,6 +71,8 @@ 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.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ + self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ self.trans_layers_weight = [ self.transformer_weight_class( mtp_layer_index, @@ -111,12 +113,7 @@ def _load_hf_weights(self): index_file = os.path.join(self.weight_dir_, "model.safetensors.index.json") assert utils.PetrelHelper.exists(index_file), "DeepSeek-V4 MTP requires model.safetensors.index.json." weight_map = utils.PetrelHelper.load_json(index_file)["weight_map"] - mtp_keys_by_file = {} - for key, file_ in weight_map.items(): - if key.startswith("mtp.0."): - mtp_keys_by_file.setdefault(file_, []).append(key) - - candidate_files = sorted(mtp_keys_by_file.keys()) + candidate_files = sorted({file_ for key, file_ in weight_map.items() if key.startswith("mtp.0.")}) assert len(candidate_files) > 0, "DeepSeek-V4 MTP weights with prefix mtp.0. were not found." loaded_key_count = 0 @@ -124,8 +121,9 @@ def _load_hf_weights(self): for file_ in tqdm(candidate_files, total=len(candidate_files), desc=desc): weights = {} with safe_open(os.path.join(self.weight_dir_, file_), "pt", "cpu") as f: - for key in mtp_keys_by_file[file_]: - weights[key] = f.get_tensor(key) + for key in f.keys(): + if key.startswith("mtp.0."): + weights[key] = f.get_tensor(key) loaded_key_count += len(weights) self.pre_post_weight.load_hf_weights(weights) From 730766cd29eef89551e1104b0e8c9e8a8e99a162 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 10:15:09 +0000 Subject: [PATCH 063/214] v4 don't need ragged_mem_buffers to save mem --- .../attention/nsa/fp8_flashmla_sparse.py | 64 ++++++++++--------- 1 file changed, 35 insertions(+), 29 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index add32918c3..c0ba1d027b 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -85,11 +85,16 @@ def _metadata_from_dict(infer_state, nsa_dict: dict) -> "_Dsv4Metadata": class NsaFlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) - device = get_current_device_id() - self.ragged_mem_buffers = [ - torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) - for _ in range(2) - ] + self.use_dsv4_flashmla_kvcache = model.config.get("model_type") == "deepseek_v4" + if self.use_dsv4_flashmla_kvcache: + self.ragged_mem_buffers = None + logger.info("DSV4 FlashMLA kvcache path skips generic NSA ragged decode buffers") + else: + device = get_current_device_id() + self.ragged_mem_buffers = [ + torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) + for _ in range(2) + ] self.prefill_flash_mla, self.prefill_flash_mla_supports_out = self._load_prefill_flash_mla() self.prefill_q_workspace = None self.prefill_out_workspace = None @@ -331,32 +336,33 @@ class NsaFlashMlaFp8SparseDecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: NsaFlashMlaFp8SparseAttBackend = self.backend - model = self.backend.model - use_cuda_graph = ( - self.infer_state.batch_size <= model.graph_max_batch_size - and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch - ) - - if use_cuda_graph: - self.ragged_mem_index = self.backend.ragged_mem_buffers[self.infer_state.microbatch_index] - else: - self.ragged_mem_index = torch.empty( - self.infer_state.total_token_num, - dtype=torch.int32, - device=get_current_device_id(), + if not self.backend.use_dsv4_flashmla_kvcache: + model = self.backend.model + use_cuda_graph = ( + self.infer_state.batch_size <= model.graph_max_batch_size + and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch ) - from lightllm.common.basemodel.triton_kernel.gen_nsa_ks_ke import gen_nsa_ks_ke - - self.ks, self.ke, self.lengths = gen_nsa_ks_ke( - b_seq_len=self.infer_state.b_seq_len, - 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.b_seq_len.shape[0], - ragged_mem_index=self.ragged_mem_index, - hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, - ) + if use_cuda_graph: + self.ragged_mem_index = self.backend.ragged_mem_buffers[self.infer_state.microbatch_index] + else: + self.ragged_mem_index = torch.empty( + self.infer_state.total_token_num, + dtype=torch.int32, + device=get_current_device_id(), + ) + + from lightllm.common.basemodel.triton_kernel.gen_nsa_ks_ke import gen_nsa_ks_ke + + self.ks, self.ke, self.lengths = gen_nsa_ks_ke( + b_seq_len=self.infer_state.b_seq_len, + 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.b_seq_len.shape[0], + ragged_mem_index=self.ragged_mem_index, + hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, + ) import flash_mla # one sched_meta per layer type: the lazy config locks extra-cache geometry (page size, From 9f4b53400e45a7452f0a2d5a18580156c7dbc066 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 10:28:15 +0000 Subject: [PATCH 064/214] fix error alloc --- .../server/router/model_infer/infer_batch.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index fa577bd5ba..d5014969d7 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1028,12 +1028,20 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: if seq_len <= 0: return 0, 0, 0 - swa_page_num = 1 if (seq_len - 1) % self.dsv4_swa_page_size == 0 else 0 + swa_page_num = 0 c4_page_num = 0 - if seq_len % 4 == 0: - entry = seq_len // 4 - 1 - c4_page_num = 1 if entry % self.dsv4_c4_page_size == 0 else 0 - c128_slot_num = 1 if self.dsv4_has_c128 and seq_len % 128 == 0 else 0 + c128_slot_num = 0 + # MTP decode prepares slots for the current token plus draft-verify rows. + for step in range(self.mtp_step + 1): + cur_seq_len = seq_len + step + if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: + swa_page_num += 1 + if cur_seq_len % 4 == 0: + entry = cur_seq_len // 4 - 1 + if entry % self.dsv4_c4_page_size == 0: + c4_page_num += 1 + if self.dsv4_has_c128 and cur_seq_len % 128 == 0: + c128_slot_num += 1 return swa_page_num, c4_page_num, c128_slot_num From a3676bb28fac39fc66f58c12fa3c312efac6fdbc Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 10:58:45 +0000 Subject: [PATCH 065/214] fix --- .../server/router/model_infer/infer_batch.py | 29 ++++++++----------- 1 file changed, 12 insertions(+), 17 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index d5014969d7..c639417d21 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -386,26 +386,21 @@ def recover_paused_reqs( if prefill_need_token_num > can_alloc_token_num: break - if can_alloc_dsv4_swa_page_num is not None: - swa_page_num = self.get_dsv4_swa_prefill_need_page_num(req, is_chuncked_prefill=False) - if swa_page_num > can_alloc_dsv4_swa_page_num: + swa_page_num = c4_page_num = c128_slot_num = 0 + if ( + can_alloc_dsv4_swa_page_num is not None + or can_alloc_dsv4_c4_page_num is not None + or can_alloc_dsv4_c128_slot_num is not None + ): + swa_page_num, c4_page_num, c128_slot_num = req.get_dsv4_prefill_need_page_and_slot_num( + is_chuncked_prefill=False + ) + if can_alloc_dsv4_swa_page_num is not None and swa_page_num > can_alloc_dsv4_swa_page_num: break - else: - swa_page_num = 0 - - if can_alloc_dsv4_c4_page_num is not None: - c4_page_num = self.get_dsv4_c4_prefill_need_page_num(req, is_chuncked_prefill=False) - if c4_page_num > can_alloc_dsv4_c4_page_num: + if can_alloc_dsv4_c4_page_num is not None and c4_page_num > can_alloc_dsv4_c4_page_num: break - else: - c4_page_num = 0 - - if can_alloc_dsv4_c128_slot_num is not None: - c128_slot_num = self.get_dsv4_c128_prefill_need_slot_num(req, is_chuncked_prefill=False) - if c128_slot_num > can_alloc_dsv4_c128_slot_num: + if can_alloc_dsv4_c128_slot_num is not None and c128_slot_num > can_alloc_dsv4_c128_slot_num: break - else: - c128_slot_num = 0 if g_infer_context.is_linear_att_mixed_model: req._linear_match_radix_cache() From f65915d04b3b894263987c1106e554ea1f7f25e5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 11:01:58 +0000 Subject: [PATCH 066/214] free infer_state.mtp_draft_input_hiddens --- lightllm/common/basemodel/basemodel.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 2f4a814e0c..6b19a9b239 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -650,6 +650,7 @@ def _decode( def _context_forward(self, infer_state: InferStateInfo): input_embs = self.pre_infer.context_forward(infer_state.input_ids, infer_state, self.pre_post_weight) + infer_state.mtp_draft_input_hiddens = None if self.args.enable_dp_prefill_balance: assert not self.args.enable_prefill_cudagraph, "not support now" infer_state.prepare_prefill_dp_balance() @@ -720,6 +721,7 @@ def _token_forward(self, infer_state: InferStateInfo): 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) + infer_state.mtp_draft_input_hiddens = None input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) for i in range(self.layers_num): From 6bab228000b985f90233b67473a2bce976a35fc1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 11:07:35 +0000 Subject: [PATCH 067/214] fix --- .../server/router/model_infer/mode_backend/base_backend.py | 3 +++ 1 file changed, 3 insertions(+) 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 6096f86988..f08ba97310 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -605,6 +605,9 @@ def _get_classed_reqs( can_alloc_token_num = g_infer_context.get_can_alloc_token_num() is_deepseek_v4 = self.is_deepseek_v4 + can_alloc_dsv4_swa_page_num = None + can_alloc_dsv4_c4_page_num = None + can_alloc_dsv4_c128_slot_num = None if is_deepseek_v4: ( can_alloc_dsv4_swa_page_num, From 6f03e722e913c1746e2e5a859b6139cda16f0bc4 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 8 Jul 2026 13:29:16 +0000 Subject: [PATCH 068/214] support dpep --- .../meta_weights/fused_moe/impl/deepgemm_impl.py | 4 +++- .../fused_moe/grouped_fused_moe_ep.py | 14 +++++++++++--- .../fused_moe/moe_silu_and_mul_mix_quant_ep.py | 8 ++++++++ .../layer_infer/transformer_layer_infer.py | 12 ++++++++++-- lightllm/server/api_start.py | 11 ----------- 5 files changed, 32 insertions(+), 17 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 72acf2430a..94709c5102 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -78,7 +78,6 @@ def _fused_experts( is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, ): - assert clamp_limit is None, "EP deepgemm fused MoE does not support clamp_limit yet" output = fused_experts( hidden_states=input_tensor, w13=w13, @@ -89,6 +88,7 @@ def _fused_experts( quant_method=self.quant_method, is_prefill=is_prefill, previous_event=None, # for overlap + clamp_limit=clamp_limit, ) return output @@ -194,6 +194,7 @@ def masked_group_gemm( masked_m: torch.Tensor, dtype: torch.dtype, expected_m: int, + clamp_limit: Optional[float] = None, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -206,6 +207,7 @@ def masked_group_gemm( w2_weight, w2_scale, expected_m=expected_m, + clamp_limit=clamp_limit, ) def prefilled_group_gemm( diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 28fe6e4304..721aaf2b25 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -66,6 +66,7 @@ def masked_group_gemm( w2: torch.Tensor, w2_scale: torch.Tensor, expected_m: int, + clamp_limit: Optional[float] = None, ): padded_m = recv_x[0].shape[1] E, N, _ = w1.shape @@ -80,7 +81,7 @@ def masked_group_gemm( _deepgemm_grouped_fp8_nt_masked(recv_x, (w1, w1_scale), gemm_out_a, masked_m, expected_m) - silu_and_mul_masked_post_quant_fwd(gemm_out_a, qsilu_out, qsilu_out_scale, block_size, masked_m) + silu_and_mul_masked_post_quant_fwd(gemm_out_a, qsilu_out, qsilu_out_scale, block_size, masked_m, limit=clamp_limit) _deepgemm_grouped_fp8_nt_masked((qsilu_out, qsilu_out_scale), (w2, w2_scale), gemm_out_b, masked_m, expected_m) return gemm_out_b @@ -193,9 +194,12 @@ def fused_experts( quant_method: Any, is_prefill: Optional[bool], previous_event: Optional[Any] = None, + clamp_limit: Optional[float] = None, ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): + if clamp_limit is not None: + raise RuntimeError("SM100 Mega MoE does not support clamped SwiGLU yet.") return mega_moe_impl(hidden_states, w13, w2, topk_weights, topk_idx, quant_method) buffer = dist_group_manager.ep_buffer if is_prefill else dist_group_manager.ep_low_latency_buffer @@ -214,6 +218,7 @@ def fused_experts( w1_scale=w13.weight_scale, w2_scale=w2.weight_scale, previous_event=previous_event, + clamp_limit=clamp_limit, ) @@ -232,6 +237,7 @@ def fused_experts_impl( w1_scale: Optional[torch.Tensor] = None, w2_scale: Optional[torch.Tensor] = None, previous_event: Optional[Any] = None, + clamp_limit: Optional[float] = None, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -321,7 +327,7 @@ def fused_experts_impl( # TODO fused kernel silu_out = torch.empty((all_tokens, N // 2), device=hidden_states.device, dtype=hidden_states.dtype) - silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out) + silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out, limit=clamp_limit) qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( silu_out, block_size_k, dtype=w1.dtype, column_major_scales=True, scale_tma_aligned=True ) @@ -365,7 +371,9 @@ def fused_experts_impl( return_recv_hook=False, ) # deepgemm - gemm_out_b = masked_group_gemm(recv_x, masked_m, hidden_states.dtype, w1, w1_scale, w2, w2_scale, expected_m) + gemm_out_b = masked_group_gemm( + recv_x, masked_m, hidden_states.dtype, w1, w1_scale, w2, w2_scale, expected_m, clamp_limit=clamp_limit + ) # low latency combine combined_x, event_overlap, hook = buffer.low_latency_combine( gemm_out_b, topk_idx, topk_weights, handle, async_finish=False, return_recv_hook=False diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py index aa91f15ed9..3c5c0de1de 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py @@ -24,8 +24,10 @@ def _silu_and_mul_post_quant_kernel( size_n, fp8_max, fp8_min, + limit: tl.constexpr, BLOCK_N: tl.constexpr, NUM_STAGE: tl.constexpr, + USE_LIMIT_ONLY: tl.constexpr = False, USE_TANH_APPROXIMATE_GELU: tl.constexpr = False, ): expert_id = tl.program_id(2) @@ -51,6 +53,9 @@ def _silu_and_mul_post_quant_kernel( for token_index in tl.range(token_id, token_num_cur_expert, block_num_per_expert, num_stages=NUM_STAGE): gate = tl.load(input_ptr_offs + token_index * stride_input_1, mask=offs_in_d < size_n, other=0.0).to(tl.float32) up = tl.load(input_ptr_offs + token_index * stride_input_1 + size_n, mask=offs_in_d < size_n, other=0.0) + if USE_LIMIT_ONLY: + gate = tl.minimum(gate, limit) + up = tl.minimum(tl.maximum(up, -limit), limit) if USE_TANH_APPROXIMATE_GELU: gate_cubed = gate * gate * gate tanh_arg = 0.7978845608028654 * (gate + 0.044715 * gate_cubed) @@ -80,6 +85,7 @@ def silu_and_mul_masked_post_quant_fwd( output_scale: torch.Tensor, quant_group_size: int, masked_m: torch.Tensor, + limit=None, ): """ input shape [expert_num, token_num_padded, hidden_dim] @@ -133,8 +139,10 @@ def silu_and_mul_masked_post_quant_fwd( size_n, fp8_max, fp8_min, + limit=limit, BLOCK_N=BLOCK_N, NUM_STAGE=NUM_STAGES, + USE_LIMIT_ONLY=limit is not None, USE_TANH_APPROXIMATE_GELU=ffn_use_tanh_approximate_gelu(), num_warps=num_warps, ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index d67ad11165..12d6ed55fe 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -340,11 +340,19 @@ def _token_attention_kernel( ) # ------------------------------------------------------------------ moe - def _routed_experts(self, x, weights, indices, layer_weight: DeepseekV4TransformerLayerWeight): + def _routed_experts( + self, + x, + weights, + indices, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + ): return layer_weight.experts_.experts_with_preselected( input_tensor=x, topk_weights=weights, topk_ids=indices, + is_prefill=infer_state.is_prefill, clamp_limit=float(self.swiglu_limit), ) @@ -371,7 +379,7 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) - routed = self._routed_experts(x, weights, indices, layer_weight) + routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) if self.enable_ep_moe: if self.tp_world_size_ > 1: all_reduce( diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index a5d622ff74..bf87c845c2 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -114,17 +114,6 @@ def normal_or_p_d_start(args): else: args.enable_multimodal = True - model_type = get_model_type(args.model_dir) - if model_type == "deepseek_v4": - if args.run_mode != "normal": - raise NotImplementedError("DeepSeek-V4 currently supports only run_mode=normal in LightLLM.") - if args.enable_cpu_cache or args.enable_disk_cache: - raise NotImplementedError("DeepSeek-V4 CPU/disk KV cache is not supported yet.") - if args.enable_ep_moe: - raise NotImplementedError("DeepSeek-V4 EP MoE is not supported yet; use TP for now.") - if "prompt_cache_kv_buffer" in get_config_json(args.model_dir): - raise NotImplementedError("DeepSeek-V4 prompt_cache_kv_buffer is not supported yet.") - if args.enable_cpu_cache: # 生成一个用于创建cpu kv cache的共享内存id。 args.cpu_kv_cache_shm_id = uuid.uuid1().int % 123456789 From c97c086f32db489550c6dba675f10d41410d8a3e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 9 Jul 2026 04:35:23 +0000 Subject: [PATCH 069/214] draft need only swa --- lightllm/common/basemodel/basemodel.py | 7 +++++-- lightllm/common/req_manager.py | 17 ++++++++++------- lightllm/models/deepseek_v4_mtp/model.py | 1 + .../server/router/model_infer/infer_batch.py | 9 ++++++++- 4 files changed, 24 insertions(+), 10 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 6b19a9b239..a9adc77e9c 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -312,6 +312,7 @@ def _prepare_decode_slots(self, model_input: ModelInput, mem_indexes: torch.Tens model_input.b_mtp_index_cpu, mem_indexes, model_input.mtp_decode_slot_prepare_indices, + prepare_compress_slots=not self.is_mtp_draft_model, ) return @@ -602,8 +603,10 @@ def _decode( 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 + if ( + self.graph is not None + and not self.is_mtp_draft_model + and self.graph.can_run(batch_size=infer_batch_size, max_len_in_batch=model_input.max_kv_seq_len) ): infer_batch_size = self.graph.find_closest_graph_batch_size(batch_size=infer_batch_size) model_input = self._create_padded_decode_model_input( diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 45ceef9013..f070a9b538 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -181,6 +181,7 @@ def prepare_decode( b_mtp_index_cpu, mem_indexes, mtp_decode_slot_prepare_indices, + prepare_compress_slots=True, ): """每个 decode step 在 attention metadata 构建前调用的钩子。基类 no-op; 需要 per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" @@ -505,9 +506,10 @@ def prepare_decode( b_mtp_index_cpu, mem_indexes, mtp_decode_slot_prepare_indices, + prepare_compress_slots=True, ): - """decode 每步槽位 prep: 先 swa 再 compress。由 BaseModel 在 copy_kv_index_to_req - 之后、attention metadata 构建前调用。""" + """decode 每步槽位 prep。由 BaseModel 在 copy_kv_index_to_req 之后、attention + metadata 构建前调用。DeepSeek-V4 MTP draft layer 只需要 SWA 槽位。""" max_mtp_index = int(b_mtp_index_cpu.max().item()) if mtp_decode_slot_prepare_indices is None: steps = range(max_mtp_index + 1) @@ -524,11 +526,12 @@ def prepare_decode( b_seq_len_cpu[rows], mem_indexes[rows], ) - self.prepare_decode_compress_slots( - b_req_idx_cpu[rows], - b_seq_len_cpu[rows], - mem_indexes[rows], - ) + if prepare_compress_slots: + self.prepare_decode_compress_slots( + b_req_idx_cpu[rows], + b_seq_len_cpu[rows], + mem_indexes[rows], + ) return def prepare_prefill( diff --git a/lightllm/models/deepseek_v4_mtp/model.py b/lightllm/models/deepseek_v4_mtp/model.py index 829ece2cd7..eddcd07d7c 100644 --- a/lightllm/models/deepseek_v4_mtp/model.py +++ b/lightllm/models/deepseek_v4_mtp/model.py @@ -92,6 +92,7 @@ def _init_infer_layer(self, start_layer_index=None): [len(previous_model.layers_infer) for previous_model in self.mtp_previous_draft_models] ) self.layers_infer = [self.transformer_layer_infer_class(total_pre_layers_num, network_config=self.config)] + assert self.layers_infer[0].compress_ratio == 0, "DeepSeek-V4 MTP draft layer must be SWA-only" return def _init_some_value(self): diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index c639417d21..fe283fd2f6 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1026,7 +1026,7 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: swa_page_num = 0 c4_page_num = 0 c128_slot_num = 0 - # MTP decode prepares slots for the current token plus draft-verify rows. + # Main model prepares current token plus draft-verify rows: SWA + compressed slots. for step in range(self.mtp_step + 1): cur_seq_len = seq_len + step if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: @@ -1037,6 +1037,13 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: c4_page_num += 1 if self.dsv4_has_c128 and cur_seq_len % 128 == 0: c128_slot_num += 1 + + # EAGLE draft forwards after the first one consume newly appended draft-only rows. + # The DeepSeek-V4 MTP draft layer is compress_ratio=0, so these rows need only SWA. + for step in range(self.mtp_step + 1, self.mtp_step * 2): + cur_seq_len = seq_len + step + if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: + swa_page_num += 1 return swa_page_num, c4_page_num, c128_slot_num From 444ca144d78e379ff95593aa4d3a70570928543b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 9 Jul 2026 14:24:56 +0000 Subject: [PATCH 070/214] speed up --- lightllm/common/req_manager.py | 73 +++++++++++++------ .../router/dynamic_prompt/radix_cache.py | 73 +++++++++++++------ .../server/router/model_infer/infer_batch.py | 1 + 3 files changed, 100 insertions(+), 47 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index f070a9b538..5fdf2e62ce 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -30,7 +30,7 @@ @dataclass class DeepseekV4PromptCachePayload: - """prompt cache 载荷: 只剩 swa 按页有效性 bitmap。 + """prompt cache 载荷: swa 按页有效性 bitmap 和最后有效页。 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4/c128 compressor 状态以 swa @@ -43,6 +43,20 @@ class DeepseekV4PromptCachePayload: cache_len: int swa_page_valid: Optional[torch.Tensor] = None + swa_last_valid_page: int = -1 + + def refresh_swa_last_valid_page(self) -> None: + if self.swa_page_valid is None: + self.swa_last_valid_page = -1 + return + valid_idx = torch.nonzero(self.swa_page_valid).flatten() + self.swa_last_valid_page = -1 if valid_idx.numel() == 0 else int(valid_idx[-1].item()) + return + + def valid_match_length(self, natural_len: int, page: int) -> int: + if self.swa_last_valid_page < 0: + return 0 + return (int(self.swa_last_valid_page) + 1) * page class DeepseekV4PromptCacheValueOps: @@ -59,25 +73,14 @@ def free(self, payload: DeepseekV4PromptCachePayload): # 槽位资源全部由 mem_manager.free(full_slots) 级联回收,载荷本身没有需要释放的资源。 return - def invalidate_swa_pages(self, payload: DeepseekV4PromptCachePayload) -> None: - """swa 压力阀回收了该节点的 swa 页后清 bitmap: 后续命中按缩短语义裁剪,不会复活。""" - if payload is not None and payload.swa_page_valid is not None: - payload.swa_page_valid.fill_(False) - return - def valid_match_length(self, payload: Optional[DeepseekV4PromptCachePayload], natural_len: int) -> int: """radix 匹配裁剪: 返回 <= natural_len 的最大 prompt-cache 边界 L',使结尾页有效。 - 有效性可能非单调(owner 生前从左驱逐、后续阀从尾回收),按候选边界回查 bitmap; - 中段 invalid 页不挡更靠后的有效命中(注意力只回看最后一个窗口)。""" - page = self.req_manager.get_prompt_cache_page_size() - if payload is None or payload.swa_page_valid is None: + 有效性可能非单调(owner 生前从左驱逐、后续阀从尾回收),中段 invalid 页不挡更 + 靠后的有效命中(注意力只回看最后一个窗口)。""" + if payload is None: return 0 - n_pages = min(natural_len // page, int(payload.swa_page_valid.numel())) - valid_idx = torch.nonzero(payload.swa_page_valid[:n_pages]) - if valid_idx.numel() == 0: - return 0 - return (int(valid_idx[-1]) + 1) * page + return payload.valid_match_length(natural_len, self.req_manager.get_prompt_cache_page_size()) class _ReqNode: @@ -451,6 +454,15 @@ def _swa_retain_len(self) -> int: V4 prompt-cache 页取 256 token,正好覆盖一个 c4 物理页对应的 token 范围。""" return int(self.sliding_window) + self.get_prompt_cache_page_size() + def _align_swa_evict_frontier(self, raw_frontier: int) -> int: + """SWA 回收水位线按 prompt-cache 页向下对齐。 + + bitmap 的有效性是 prompt-cache page 粒度;若水位线切进页面中间,该页会被判为 + invalid,即使靠近命中边界的窗口实际仍完整驻留。""" + page = self.get_prompt_cache_page_size() + raw_frontier = max(0, int(raw_frontier)) + return raw_frontier // page * page + def prepare_prefill_swa( self, b_req_idx: torch.Tensor, @@ -480,9 +492,9 @@ def prepare_prefill_swa( mark = self._swa_evict_marks[req_idx] if mark < 0: # 首个 chunk: [0, ready_len) 是 radix 共享前缀,其 swa 槽归 radix 所有,不可回收。 - self._swa_evict_marks[req_idx] = ready_len + self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(ready_len) continue - evict_end = ready_len - retain + 1 + evict_end = self._align_swa_evict_frontier(ready_len - retain + 1) if evict_end > mark: evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) self._swa_evict_marks[req_idx] = evict_end @@ -585,9 +597,9 @@ def prepare_decode_swa( mark = self._swa_evict_marks[req_idx] if mark < 0: # 未经过 prefill prep 的保守路径: 不回收旧位置,仅推进水位线。 - self._swa_evict_marks[req_idx] = max(0, seq_len - retain) + self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(seq_len - retain) continue - evict_end = seq_len - retain + evict_end = self._align_swa_evict_frontier(seq_len - retain) if evict_end > mark: evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) self._swa_evict_marks[req_idx] = evict_end @@ -955,7 +967,7 @@ def compute_swa_page_valid(self, full_slots: torch.Tensor) -> torch.Tensor: def swa_page_valid_from_watermark(self, req_idx: int, cache_len: int) -> torch.Tensor: """插入时的按页有效性,纯 CPU: 请求自有 token 的 swa 映射只被出窗水位线回收 - (阀不触活跃请求,级联只在 free 时),页 p 全驻留 ⟺ 页起点 128p >= 水位线。 + (阀不触活跃请求,级联只在 free 时),页 p 全驻留 ⟺ 页起点 page*p >= 水位线。 与 compute_swa_page_valid 在插入时刻对自有 token 等价,但不做 GPU gather/同步—— router 关键路径上每次插入省一次对全部在途 kernel 的等待。bitmap 中借入前缀 @@ -971,21 +983,36 @@ def slice_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload, star end = int(end) page = self.get_prompt_cache_page_size() # radix page 保证分裂点页对齐,bitmap 可整页切分。 - return DeepseekV4PromptCachePayload( + ans = DeepseekV4PromptCachePayload( cache_len=end - start, swa_page_valid=payload.swa_page_valid[start // page : end // page].clone() if payload.swa_page_valid is not None else None, ) + ans.refresh_swa_last_valid_page() + return ans def concat_prompt_cache_payloads(self, payloads: List[DeepseekV4PromptCachePayload]): if len(payloads) == 0: return None bitmaps = [p.swa_page_valid for p in payloads] - return DeepseekV4PromptCachePayload( + ans = DeepseekV4PromptCachePayload( cache_len=sum(p.cache_len for p in payloads), swa_page_valid=torch.cat(bitmaps, dim=0) if all(b is not None for b in bitmaps) else None, ) + if ans.swa_page_valid is None: + return ans + + page = self.get_prompt_cache_page_size() + page_offset = 0 + last_valid_page = -1 + for item in payloads: + item_last = int(getattr(item, "swa_last_valid_page", -1)) + if item_last >= 0: + last_valid_page = page_offset + item_last + page_offset += int(item.cache_len) // page + ans.swa_last_valid_page = last_valid_page + return ans def build_prompt_cache_payload( self, diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index e45b556b57..1d58cb8dea 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -660,35 +660,60 @@ def _print_helper(self, node: TreeNode, indent): return def free_unreferenced_swa_pages(self, need_pages: int) -> None: - """DeepSeek-V4 swa free hook: 页 allocator 触底时,沿 LRU 序(evict_tree_set)只对 - ref_count==0 的节点链 free 其 swa 页(full 槽与压缩条目保留——节点仍可服务更长前缀的 - 中段命中),并清载荷 bitmap 位使后续命中按缩短语义裁剪。所有权判定直接复用 radix - 引用计数: 节点被任何活跃请求借用即 ref>0,其页不可达。不够时由 allocator 的 assert - 兜底(最后防线)。""" + """DeepSeek-V4 swa free hook: 页 allocator 触底时,回收 ref_count==0 节点的 swa 页。""" if self.mem_manager is None or self.extra_value_ops is None: return - invalidate = getattr(self.extra_value_ops, "invalidate_swa_pages", None) - if invalidate is None: - return allocator = self.mem_manager.swa_page_allocator target = allocator.can_use_mem_size + int(need_pages) - for leaf in list(self.evict_tree_set): - if allocator.can_use_mem_size >= target: + evict_slots = [] + invalidate_payloads = [] + evict_swa_pages = 0 + for free_last in (False, True): + visited = set() + for leaf in self.evict_tree_set: + if allocator.can_use_mem_size + evict_swa_pages >= target: + break + node = leaf + while node is not None and node is not self.root_node and node.ref_counter == 0: + node_id = id(node) + if node_id in visited: + node = node.parent + continue + visited.add(node_id) + + payload = node.token_extra_value + if ( + len(node.token_mem_index_value) > 0 + and payload is not None + and payload.swa_page_valid is not None + ): + last_page = int(payload.swa_last_valid_page) + if last_page >= 0: + if free_last: + page_slice = slice(last_page, last_page + 1) + else: + page_slice = slice(0, last_page) + valid_pages = int(payload.swa_page_valid[page_slice].sum().item()) + if valid_pages > 0: + start = page_slice.start * self.page_size + end = min(page_slice.stop * self.page_size, len(node.token_mem_index_value)) + if end > start: + evict_slots.append(node.token_mem_index_value[start:end]) + invalidate_payloads.append((payload, page_slice, free_last)) + evict_swa_pages += valid_pages * self._swa_pages_per_prompt_page + if allocator.can_use_mem_size + evict_swa_pages >= target: + break + node = node.parent + if allocator.can_use_mem_size + evict_swa_pages >= target: break - node = leaf - # 叶子起步沿父链回收: 引用计数向上累加(add_node_ref_counter 走父链), - # 因此 ref==0 的祖先必无任何活跃借用方。重复访问无害(evict_swa/-1 跳过)。 - # 每回收一个节点就复查目标,避免多回收(无谓削减命中可用性)。 - while node is not None and node is not self.root_node and node.ref_counter == 0: - if len(node.token_mem_index_value) > 0: - old_pages = self._node_swa_pages_num(node) - self.mem_manager.evict_swa(node.token_mem_index_value) - if node.token_extra_value is not None: - invalidate(node.token_extra_value) - self.swa_tree_total_pages_num -= old_pages - self._node_swa_pages_num(node) - if allocator.can_use_mem_size >= target: - return - node = node.parent + if len(evict_slots) == 0: + return + self.mem_manager.evict_swa(torch.cat(evict_slots)) + for payload, page_slice, free_last in invalidate_payloads: + payload.swa_page_valid[page_slice] = False + if free_last: + payload.swa_last_valid_page = -1 + self.swa_tree_total_pages_num -= evict_swa_pages return def free_radix_cache_to_get_enough_token(self, need_token_num): diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index fe283fd2f6..6dfc35dbc2 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -173,6 +173,7 @@ def _dsv4_full_att_free_req(self, free_token_index: List, req: "InferReq"): value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() payload.swa_page_valid = self.req_manager.swa_page_valid_from_watermark(req.req_idx, cache_len) + payload.refresh_swa_last_valid_page() key = torch.tensor(req.get_input_token_ids()[0:cache_len], dtype=torch.int64, device="cpu") duplicate_prefix_len, _ = self.radix_cache.insert(key, value[:cache_len], extra_value=payload) From a0acc0610ee714e3d5ffc255c3848310f2070a25 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 01:16:34 +0000 Subject: [PATCH 071/214] support overlap --- lightllm/server/router/model_infer/mode_backend/base_backend.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index f08ba97310..16460bf351 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -155,8 +155,6 @@ def init_model(self, kvargs): set_random_seed(2147483647) self.is_linear_att_mixed_model = isinstance(self.model.req_manager, ReqManagerForMamba) self.is_deepseek_v4 = isinstance(self.model.req_manager, DeepseekV4ReqManager) - if self.is_deepseek_v4: - self.support_overlap = False if self.is_linear_att_mixed_model: self.linear_att_cache_manager = LinearAttCacheManager( From b8c3a63b35139cf771e817ef680b138a94e9e6e4 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 01:58:30 +0000 Subject: [PATCH 072/214] refact --- .../fused_moe/fused_moe_weight.py | 19 ++++++------------- .../layer_infer/transformer_layer_infer.py | 2 +- 2 files changed, 7 insertions(+), 14 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 24842ed383..3e11e39852 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -137,7 +137,6 @@ def experts( is_prefill: Optional[bool] = None, ) -> torch.Tensor: """Backward compatible method that routes to platform-specific implementation.""" - self._finalize_moe_weight() return self.fuse_moe_impl( input_tensor=input_tensor, router_logits=router_logits, @@ -154,7 +153,7 @@ def experts( per_expert_scale=self.per_expert_scale, ) - def experts_with_preselected( + def experts_with_topk( self, input_tensor: torch.Tensor, topk_weights: torch.Tensor, @@ -162,7 +161,6 @@ def experts_with_preselected( is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, ) -> torch.Tensor: - self._finalize_moe_weight() return self.fuse_moe_impl.fused_experts_with_topk( input_tensor=input_tensor, w13=self.w13, @@ -302,18 +300,13 @@ def verify_load(self): True if self.e_score_correction_bias is None else getattr(self.e_score_correction_bias, "load_ok", False) ) load_ok = weight_load_ok and per_expert_scale_load_ok and e_score_correction_bias_load_ok - if load_ok: - self._finalize_moe_weight() + if load_ok and not self._moe_weight_finalized: + finalize = getattr(self.quant_method, "finalize_moe_weight", None) + if finalize is not None: + finalize(self) + self._moe_weight_finalized = True return load_ok - def _finalize_moe_weight(self): - if self._moe_weight_finalized: - return - finalize = getattr(self.quant_method, "finalize_moe_weight", None) - if finalize is not None: - finalize(self) - self._moe_weight_finalized = True - def _create_weight(self): intermediate_size = self.split_inter_size self.e_score_correction_bias = None diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 12d6ed55fe..3cd3ec7010 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -348,7 +348,7 @@ def _routed_experts( infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, ): - return layer_weight.experts_.experts_with_preselected( + return layer_weight.experts_.experts_with_topk( input_tensor=x, topk_weights=weights, topk_ids=indices, From 1a2f29d05ec2ae90fe608ff69e97995df40d6ba6 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 02:55:43 +0000 Subject: [PATCH 073/214] mega_moe support clamp_limit --- .../triton_kernel/fused_moe/grouped_fused_moe_ep.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 721aaf2b25..c5b241f1ab 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -120,6 +120,7 @@ def mega_moe_impl( topk_weights: torch.Tensor, topk_ids: torch.Tensor, quant_method: Any, + clamp_limit: Optional[float] = None, ): if not (HAS_DEEPGEMM and hasattr(deep_gemm, "fp8_fp4_mega_moe")): raise RuntimeError("deep_gemm does not provide fp8-fp4 Mega MoE kernel") @@ -157,6 +158,7 @@ def mega_moe_impl( l2_weights, buffer, cumulative_local_expert_recv_stats=stats, + activation_clamp=clamp_limit, ) return output @@ -198,9 +200,7 @@ def fused_experts( ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): - if clamp_limit is not None: - raise RuntimeError("SM100 Mega MoE does not support clamped SwiGLU yet.") - return mega_moe_impl(hidden_states, w13, w2, topk_weights, topk_idx, quant_method) + return mega_moe_impl(hidden_states, w13, w2, topk_weights, topk_idx, quant_method, clamp_limit=clamp_limit) buffer = dist_group_manager.ep_buffer if is_prefill else dist_group_manager.ep_low_latency_buffer return fused_experts_impl( From 24084627a94d7d4d9810c607955578f2e6029fa8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 05:31:44 +0000 Subject: [PATCH 074/214] Move DSV4 SWA/c4/c128 slot preparation into the model-specific path. Derive current and MTP predecessor slots directly from mem_indexes, and remove BaseModel hooks, CPU mirror padding, and redundant CUDA casts. --- lightllm/common/basemodel/basemodel.py | 73 ---- lightllm/common/basemodel/batch_objs.py | 2 - .../deepseek4_mem_manager.py | 101 ++--- lightllm/common/req_manager.py | 399 +++++++----------- lightllm/models/deepseek_v4/model.py | 48 +++ lightllm/server/api_start.py | 2 - 6 files changed, 238 insertions(+), 387 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index a9adc77e9c..9f986fc873 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -301,21 +301,6 @@ def forward(self, model_input: ModelInput): else: return self._decode(model_input) - def _prepare_decode_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor): - if model_input.mem_indexes_cpu is None: - return - if model_input.mtp_decode_slot_prepare_indices == (): - return - self.req_manager.prepare_decode( - model_input.b_req_idx_cpu, - model_input.b_seq_len_cpu, - model_input.b_mtp_index_cpu, - mem_indexes, - model_input.mtp_decode_slot_prepare_indices, - prepare_compress_slots=not self.is_mtp_draft_model, - ) - return - def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() infer_state.input_ids = model_input.input_ids @@ -387,19 +372,7 @@ def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_s new_model_input.b_mtp_index = F.pad( new_model_input.b_mtp_index, (0, padded_batch_size), mode="constant", value=0 ) - new_model_input.b_mtp_index_cpu = F.pad( - new_model_input.b_mtp_index_cpu, (0, padded_batch_size), mode="constant", value=0 - ) new_model_input.b_seq_len = F.pad(new_model_input.b_seq_len, (0, padded_batch_size), mode="constant", value=2) - new_model_input.b_req_idx_cpu = F.pad( - new_model_input.b_req_idx_cpu, - (0, padded_batch_size), - mode="constant", - value=self.req_manager.HOLD_REQUEST_ID, - ) - new_model_input.b_seq_len_cpu = F.pad( - new_model_input.b_seq_len_cpu, (0, padded_batch_size), mode="constant", value=2 - ) new_model_input.mem_indexes = F.pad( new_model_input.mem_indexes, (0, padded_batch_size), @@ -455,18 +428,8 @@ def _create_padded_prefill_model_input(self, model_input: ModelInput, new_handle new_model_input.b_req_idx, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID ) new_model_input.b_mtp_index = F.pad(new_model_input.b_mtp_index, (0, 1), mode="constant", value=0) - new_model_input.b_mtp_index_cpu = F.pad(new_model_input.b_mtp_index_cpu, (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_req_idx_cpu = F.pad( - new_model_input.b_req_idx_cpu, (0, 1), mode="constant", value=self.req_manager.HOLD_REQUEST_ID - ) - new_model_input.b_seq_len_cpu = F.pad( - new_model_input.b_seq_len_cpu, (0, 1), mode="constant", value=padded_token_num - ) - new_model_input.b_ready_cache_len_cpu = F.pad( - new_model_input.b_ready_cache_len_cpu, (0, 1), mode="constant", value=0 - ) 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 可能会让外面使用的数组引用发生变化,导致错误。 @@ -560,15 +523,6 @@ def _prefill( alloc_mem_index=infer_state.mem_index, max_q_seq_len=infer_state.max_q_seq_len, ) - if model_input.b_req_idx_cpu is not None and not self.is_mtp_draft_model: - self.req_manager.prepare_prefill( - b_req_idx=infer_state.b_req_idx, - b_ready_cache_len=infer_state.b_ready_cache_len, - b_seq_len=infer_state.b_seq_len, - b_req_idx_cpu=model_input.b_req_idx_cpu, - b_ready_cache_len_cpu=model_input.b_ready_cache_len_cpu, - b_seq_len_cpu=model_input.b_seq_len_cpu, - ) prefill_mem_indexes_ready_event = torch.cuda.Event() prefill_mem_indexes_ready_event.record() @@ -596,7 +550,6 @@ def _decode( model_input.b_mtp_index, ) - origin_model_input = model_input origin_batch_size = model_input.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_ @@ -619,7 +572,6 @@ def _decode( infer_state.b_seq_len, infer_state.mem_index, ) - self._prepare_decode_slots(origin_model_input, infer_state.mem_index[:origin_batch_size]) infer_state.init_some_extra_state(self) infer_state.init_att_state() @@ -641,7 +593,6 @@ def _decode( infer_state.b_seq_len, infer_state.mem_index, ) - self._prepare_decode_slots(origin_model_input, infer_state.mem_index[:origin_batch_size]) infer_state.init_some_extra_state(self) infer_state.init_att_state() model_output = self._token_forward(infer_state) @@ -808,15 +759,6 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state0.mem_index, max_q_seq_len=infer_state0.max_q_seq_len, ) - if model_input0.b_req_idx_cpu is not None and not self.is_mtp_draft_model: - self.req_manager.prepare_prefill( - b_req_idx=infer_state0.b_req_idx, - b_ready_cache_len=infer_state0.b_ready_cache_len, - b_seq_len=infer_state0.b_seq_len, - b_req_idx_cpu=model_input0.b_req_idx_cpu, - b_ready_cache_len_cpu=model_input0.b_ready_cache_len_cpu, - b_seq_len_cpu=model_input0.b_seq_len_cpu, - ) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -830,15 +772,6 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod alloc_mem_index=infer_state1.mem_index, max_q_seq_len=infer_state1.max_q_seq_len, ) - if model_input1.b_req_idx_cpu is not None and not self.is_mtp_draft_model: - self.req_manager.prepare_prefill( - b_req_idx=infer_state1.b_req_idx, - b_ready_cache_len=infer_state1.b_ready_cache_len, - b_seq_len=infer_state1.b_seq_len, - b_req_idx_cpu=model_input1.b_req_idx_cpu, - b_ready_cache_len_cpu=model_input1.b_ready_cache_len_cpu, - b_seq_len_cpu=model_input1.b_seq_len_cpu, - ) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -888,8 +821,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode assert model_input1.mem_indexes.is_cuda origin_batch_size = model_input0.batch_size - origin_model_input0 = model_input0 - origin_model_input1 = model_input1 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_ @@ -906,7 +837,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.b_seq_len, infer_state0.mem_index, ) - self._prepare_decode_slots(origin_model_input0, infer_state0.mem_index[:origin_batch_size]) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -917,7 +847,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.b_seq_len, infer_state1.mem_index, ) - self._prepare_decode_slots(origin_model_input1, infer_state1.mem_index[:origin_batch_size]) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() @@ -949,7 +878,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state0.b_seq_len, infer_state0.mem_index, ) - self._prepare_decode_slots(origin_model_input0, infer_state0.mem_index[:origin_batch_size]) infer_state0.init_some_extra_state(self) infer_state0.init_att_state() @@ -960,7 +888,6 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode infer_state1.b_seq_len, infer_state1.mem_index, ) - self._prepare_decode_slots(origin_model_input1, infer_state1.mem_index[:origin_batch_size]) infer_state1.init_some_extra_state(self) infer_state1.init_att_state() diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 7a4d837081..9e56beeb4c 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -46,7 +46,6 @@ class ModelInput: b_req_idx_cpu: torch.Tensor = None b_mtp_index_cpu: torch.Tensor = None b_seq_len_cpu: torch.Tensor = None - b_ready_cache_len_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 # 的一些变量 b_prefill_has_output_cpu: List[bool] = None # 标记进行prefill的请求是否具有输出 @@ -70,7 +69,6 @@ def capture_cpu_mirrors(self): self._capture_cpu_mirror("b_req_idx", "b_req_idx_cpu") self._capture_cpu_mirror("b_mtp_index", "b_mtp_index_cpu") self._capture_cpu_mirror("b_seq_len", "b_seq_len_cpu") - self._capture_cpu_mirror("b_ready_cache_len", "b_ready_cache_len_cpu") return def make_mtp_draft_input(self): diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index b37b18eee4..1e468a4eef 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -372,22 +372,20 @@ def _alloc_swa_pages(self, need_pages: int) -> torch.Tensor: self._free_radix_unreferenced_swa_fn(need_pages - self.swa_page_allocator.can_use_mem_size) return self.swa_page_allocator.alloc(need_pages) - def _count_swa_pages(self, swa_slots: torch.Tensor, delta: int) -> torch.Tensor: - """按 slot 所在页更新存活计数,返回触达的页(去重)。""" - pages = torch.div(swa_slots.long(), DSV4_SWA_PAGE_SIZE, rounding_mode="floor") + def _update_swa_page_counts(self, swa_slots: torch.Tensor, delta: int) -> torch.Tensor: + """按 slot 所在页更新存活计数,返回逐 slot 的页号。""" + pages = torch.div(swa_slots, DSV4_SWA_PAGE_SIZE, rounding_mode="floor") ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) self.swa_page_live_count.index_add_(0, pages, ones) - return torch.unique(pages) + return pages def alloc_swa_prefill( self, - b_req_idx: torch.Tensor, - b_ready_cache_len: torch.Tensor, - b_seq_len: torch.Tensor, + mem_indexes: torch.Tensor, req_to_token_indexs: torch.Tensor, - b_req_idx_cpu: torch.Tensor, - b_ready_cache_len_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + ready_list: List[int], + seq_list: List[int], ) -> None: """prefill prep: 为各请求位置 [ready, seq) 的新 token 分配位置对齐的 swa 槽。 @@ -395,88 +393,81 @@ def alloc_swa_prefill( 续页(start 非整页,只可能是首页)的 base 从上一 token 的映射派生 (full_to_swa[req_to_token[req, start-1]],该 token 必在保留窗内);其余页全新分配。 radix 命中(ready 必 128 对齐)的借用方从全新页开始,与节点持有页天然不相交。 - 必须在 init_req_to_token_indexes 之后调用(scatter 目标经 req_to_token 行)。 + 当前 chunk 的 full 槽直接来自 generic preprocess 分配的 mem_indexes,因此不依赖 + req_to_token_indexs 已完成当前 chunk 的 scatter;只有续页的上一 token 查询旧 req 行。 """ page = DSV4_SWA_PAGE_SIZE hold_req_id = self.max_request_num # padding 行的请求 id(req_manager.HOLD_REQUEST_ID) - req_list = b_req_idx_cpu.tolist() - ready_list = b_ready_cache_len_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() - - segs = [] # (req_idx, start, end, n_new_pages, has_cont_page) + segs = [] # (req_idx, start, end, mem_offset, n_new_pages, has_cont_page) total_new_pages = 0 + mem_offset = 0 for req_idx, start, end in zip(req_list, ready_list, seq_list): - req_idx, start, end = int(req_idx), int(start), int(end) + q_len = end - start if req_idx == hold_req_id or end <= start: + mem_offset += q_len continue first_new_page = _ceil_div(start, page) n_new = max(0, (end - 1) // page - first_new_page + 1) - segs.append((req_idx, start, end, n_new, start % page != 0)) + segs.append((req_idx, start, end, mem_offset, n_new, start % page != 0)) total_new_pages += n_new + mem_offset += q_len if not segs: return - new_pages = self._alloc_swa_pages(total_new_pages).cuda(non_blocking=True).long() if total_new_pages else None + device = self.full_to_swa_indexs.device + mem_indexes = mem_indexes.reshape(-1) + new_pages = self._alloc_swa_pages(total_new_pages).to(device, non_blocking=True) if total_new_pages else None page_cursor = 0 - for req_idx, start, end, n_new, has_cont in segs: - positions = torch.arange(start, end, dtype=torch.long, device="cuda") + for req_idx, start, end, mem_start, n_new, has_cont in segs: + positions = torch.arange(start, end, dtype=torch.int32, device=device) page_local = torch.div(positions, page, rounding_mode="floor") - start // page - bases = torch.empty(((end - 1) // page - start // page + 1,), dtype=torch.long, device="cuda") + bases = torch.empty(((end - 1) // page - start // page + 1,), dtype=torch.int32, device=device) if has_cont: - prev_slot = int(self.full_to_swa_indexs[req_to_token_indexs[req_idx, start - 1].long()].item()) - # 续页不变式: 上一 token 必驻留(retain >= 2)且位置对齐(未来 resume/MTP 改动的哨兵)。 - assert prev_slot >= 0 and prev_slot % page == (start - 1) % page + prev_slot = self.full_to_swa_indexs[req_to_token_indexs[req_idx, start - 1]] bases[0] = prev_slot - (start - 1) % page if n_new: bases[1 if has_cont else 0 :] = new_pages[page_cursor : page_cursor + n_new] * page page_cursor += n_new - slots = (bases[page_local] + positions % page).to(torch.int32) - self.full_to_swa_indexs[req_to_token_indexs[req_idx, start:end].long()] = slots - self._count_swa_pages(slots, 1) + slots = bases[page_local] + positions % page + full_slots = mem_indexes[mem_start : mem_start + end - start] + self.full_to_swa_indexs[full_slots] = slots + self._update_swa_page_counts(slots, 1) return def alloc_swa_decode( self, - b_req_idx_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + seq_list: List[int], mem_indexes: torch.Tensor, - req_to_token_indexs: torch.Tensor, + prev_full_indexes: torch.Tensor, ) -> None: """decode prep: 本步 token(位置 seq-1)的 swa 槽。整页起点开新页,否则上一 token 槽 +1 (位置对齐不变式保证同页连续)。scatter 目标用当前步 mem_indexes。 - 注意: 续槽从上一位置的映射派生,故同一请求的多行(MTP 多 token/步)需要调用方按 - b_mtp_index 分段准备。""" + 调用方传入每行前一 token 的 full 槽;MTP step>0 可直接使用同批前一列。""" page = DSV4_SWA_PAGE_SIZE hold_req_id = self.max_request_num - req_list = b_req_idx_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() - cont_rows, cont_prev_pos, new_rows = [], [], [] + cont_rows, new_rows = [], [] for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): - req_idx, seq_len = int(req_idx), int(seq_len) if req_idx == hold_req_id or seq_len <= 0: continue if (seq_len - 1) % page == 0: new_rows.append(i) else: cont_rows.append(i) - cont_prev_pos.append(seq_len - 2) - mem_indexes = mem_indexes.cuda().long().reshape(-1) + mem_indexes = mem_indexes.reshape(-1) if cont_rows: - req_rows = torch.tensor([req_list[i] for i in cont_rows], dtype=torch.long, device="cuda") - prev_full = req_to_token_indexs[req_rows, torch.tensor(cont_prev_pos, device="cuda")].long() + prev_full = prev_full_indexes.reshape(-1)[cont_rows] prev_slots = self.full_to_swa_indexs[prev_full] - # 续槽不变式哨兵: 上一位置必驻留(retain 覆盖)。prep 阶段本就有同步,代价可忽略。 - assert bool((prev_slots >= 0).all()) slots = prev_slots + 1 self.full_to_swa_indexs[mem_indexes[cont_rows]] = slots - self._count_swa_pages(slots, 1) + self._update_swa_page_counts(slots, 1) if new_rows: - pages = self._alloc_swa_pages(len(new_rows)).cuda(non_blocking=True).long() - slots = (pages * page).to(torch.int32) + pages = self._alloc_swa_pages(len(new_rows)).to(self.full_to_swa_indexs.device, non_blocking=True) + slots = pages * page self.full_to_swa_indexs[mem_indexes[new_rows]] = slots - self._count_swa_pages(slots, 1) + self._update_swa_page_counts(slots, 1) return def evict_swa(self, full_slots: torch.Tensor) -> None: @@ -484,7 +475,7 @@ def evict_swa(self, full_slots: torch.Tensor) -> None: 未映射(-1)的槽位跳过;页计数减到 0 时整页归还 allocator。""" if full_slots.numel() == 0: return - full_slots = full_slots.cuda().long().reshape(-1) + full_slots = full_slots.to(self.full_to_swa_indexs.device, non_blocking=True).reshape(-1) full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) if full_slots.numel() == 0: return @@ -494,14 +485,14 @@ def evict_swa(self, full_slots: torch.Tensor) -> None: if valid_slots.numel() == 0: return self.full_to_swa_indexs[full_slots[valid]] = -1 - touched = self._count_swa_pages(valid_slots, -1) + touched = torch.unique(self._update_swa_page_counts(valid_slots, -1)) empty = touched[self.swa_page_live_count[touched] == 0] if empty.numel() > 0: self.swa_page_allocator.free(empty.to(torch.int32)) return def _evict_compress(self, full_slots: torch.Tensor, mapping: torch.Tensor, allocator: KvCacheAllocator) -> None: - full_slots = full_slots.cuda().long().reshape(-1) + full_slots = full_slots.to(mapping.device, non_blocking=True).reshape(-1) # 去重: 同批重复槽会 gather 出重复的压缩槽 -> allocator 双重释放(free 已去重,直呼叫方防御)。 full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) if full_slots.numel() == 0: @@ -520,18 +511,18 @@ def alloc_c4_pages(self, need_pages: int) -> torch.Tensor: return self.c4_page_allocator.alloc(need_pages) def count_c4_slots(self, c4_slots: torch.Tensor, delta: int) -> torch.Tensor: - """按 c4 slot 所在页更新存活计数,返回触达的页(去重)。""" + """按 c4 slot 所在页更新存活计数,返回逐 slot 的页号。""" assert self.c4_page_live_count is not None, "DeepSeek-V4 c4 page live count is not initialized" - pages = torch.div(c4_slots.long(), DSV4_C4_PAGE_SIZE, rounding_mode="floor") + pages = torch.div(c4_slots, DSV4_C4_PAGE_SIZE, rounding_mode="floor") ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) self.c4_page_live_count.index_add_(0, pages, ones) - return torch.unique(pages) + return pages def evict_c4(self, full_slots: torch.Tensor) -> None: """回收 full 槽位(组末 token)映射的 c4 槽。非组末/未映射(-1)的槽位跳过。""" if self.c4_page_allocator is None or full_slots.numel() == 0: return - full_slots = full_slots.cuda().long().reshape(-1) + full_slots = full_slots.to(self.full_to_c4_indexs.device, non_blocking=True).reshape(-1) full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) if full_slots.numel() == 0: return @@ -541,7 +532,7 @@ def evict_c4(self, full_slots: torch.Tensor) -> None: if valid_slots.numel() == 0: return self.full_to_c4_indexs[full_slots[valid]] = -1 - touched = self.count_c4_slots(valid_slots, -1) + touched = torch.unique(self.count_c4_slots(valid_slots, -1)) empty = touched[self.c4_page_live_count[touched] == 0] if empty.numel() > 0: self.c4_page_allocator.free(empty.to(torch.int32)) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 5fdf2e62ce..fe49e9a095 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -164,32 +164,6 @@ def free_all(self): self.req_list = _ReqLinkedList(self.max_request_num) return - def prepare_prefill( - self, - b_req_idx, - b_ready_cache_len, - b_seq_len, - b_req_idx_cpu=None, - b_ready_cache_len_cpu=None, - b_seq_len_cpu=None, - ): - """prefill 在 init_req_to_token_indexes 之后调用的钩子。基类 no-op; 需要 - prefill KV 槽位 prep 的模型 (DeepSeek-V4) override。""" - return - - def prepare_decode( - self, - b_req_idx_cpu, - b_seq_len_cpu, - b_mtp_index_cpu, - mem_indexes, - mtp_decode_slot_prepare_indices, - prepare_compress_slots=True, - ): - """每个 decode step 在 attention metadata 构建前调用的钩子。基类 no-op; 需要 - per-step KV 槽位 prep 的模型 (DeepSeek-V4) override。""" - return - class ReqSamplingParamsManager: """ @@ -465,30 +439,25 @@ def _align_swa_evict_frontier(self, raw_frontier: int) -> int: def prepare_prefill_swa( self, - b_req_idx: torch.Tensor, - b_ready_cache_len: torch.Tensor, - b_seq_len: torch.Tensor, - b_req_idx_cpu: torch.Tensor, - b_ready_cache_len_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + ready_list: List[int], + seq_list: List[int], + mem_indexes: torch.Tensor, ) -> None: """prefill prep: 为本 chunk 全部新 token(位置 [ready, seq))分配位置对齐的 swa 槽, 并回收已出窗位置的槽。 本 chunk 起点 L = ready_cache_len,首个新 token(位置 L)的窗口是 [L-W+1, L];回收 边界再额外保留一个 radix 页(_swa_retain_len),即位置 < L-retain+1。先回收再分配。 - 必须在 init_req_to_token_indexes 之后调用(位置对齐分配经 req_to_token 行派生/scatter)。""" + 当前 chunk 的 full slots 直接使用 generic preprocess 分配的 mem_indexes,因而可以 + 在通用 req_to_token scatter 之前执行。""" self.mem_manager: DeepseekV4MemoryManager if self.sliding_window is not None: retain = self._swa_retain_len() evict_slots = [] - req_list = b_req_idx_cpu.tolist() - ready_list = b_ready_cache_len_cpu.tolist() for req_idx, ready_len in zip(req_list, ready_list): - req_idx = int(req_idx) if req_idx == self.HOLD_REQUEST_ID: continue - ready_len = int(ready_len) mark = self._swa_evict_marks[req_idx] if mark < 0: # 首个 chunk: [0, ready_len) 是 radix 共享前缀,其 swa 槽归 radix 所有,不可回收。 @@ -501,13 +470,11 @@ def prepare_prefill_swa( if evict_slots: self.mem_manager.evict_swa(torch.cat(evict_slots)) self.mem_manager.alloc_swa_prefill( - b_req_idx, - b_ready_cache_len, - b_seq_len, + mem_indexes, self.req_to_token_indexs, - b_req_idx_cpu=b_req_idx_cpu, - b_ready_cache_len_cpu=b_ready_cache_len_cpu, - b_seq_len_cpu=b_seq_len_cpu, + req_list=req_list, + ready_list=ready_list, + seq_list=seq_list, ) return @@ -520,8 +487,8 @@ def prepare_decode( mtp_decode_slot_prepare_indices, prepare_compress_slots=True, ): - """decode 每步槽位 prep。由 BaseModel 在 copy_kv_index_to_req 之后、attention - metadata 构建前调用。DeepSeek-V4 MTP draft layer 只需要 SWA 槽位。""" + """decode 每步槽位 prep。在 BaseModel 的通用 req scatter 与 attention metadata + 构建前调用;DeepSeek-V4 MTP draft layer 只需要 SWA 槽位。""" max_mtp_index = int(b_mtp_index_cpu.max().item()) if mtp_decode_slot_prepare_indices is None: steps = range(max_mtp_index + 1) @@ -531,55 +498,60 @@ def prepare_decode( batch_size = b_mtp_index_cpu.shape[0] slots_per_req = max_mtp_index + 1 assert batch_size % slots_per_req == 0 + req_list = b_req_idx_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() + mem_indexes_by_req = mem_indexes.reshape(-1, slots_per_req) for step in steps: - rows = slice(step, batch_size, slots_per_req) + step_req_list = req_list[step::slots_per_req] + step_seq_list = seq_list[step::slots_per_req] self.prepare_decode_swa( - b_req_idx_cpu[rows], - b_seq_len_cpu[rows], - mem_indexes[rows], + step_req_list, + step_seq_list, + mem_indexes_by_req[:, step], + prev_mem_indexes=mem_indexes_by_req[:, step - 1] if step > 0 else None, ) if prepare_compress_slots: self.prepare_decode_compress_slots( - b_req_idx_cpu[rows], - b_seq_len_cpu[rows], - mem_indexes[rows], + step_req_list, + step_seq_list, + mem_indexes_by_req[:, step], + prev_group_end_mem_indexes=mem_indexes_by_req[:, step - 4] if step >= 4 else None, ) return def prepare_prefill( self, - b_req_idx: torch.Tensor, - b_ready_cache_len: torch.Tensor, - b_seq_len: torch.Tensor, b_req_idx_cpu: torch.Tensor, b_ready_cache_len_cpu: torch.Tensor, b_seq_len_cpu: torch.Tensor, + mem_indexes: torch.Tensor, ) -> None: - """prefill 槽位 prep: 先 swa 再 compress。由 BaseModel 在 - init_req_to_token_indexes 之后、attention metadata 构建之前调用。""" + """prefill 槽位 prep: 直接消费 generic preprocess 分配的 full slots,在 + BaseModel 的通用 req scatter 与 attention metadata 构建之前完成。""" + req_list = b_req_idx_cpu.tolist() + ready_list = b_ready_cache_len_cpu.tolist() + seq_list = b_seq_len_cpu.tolist() + mem_indexes = mem_indexes.reshape(-1) self.prepare_prefill_swa( - b_req_idx, - b_ready_cache_len, - b_seq_len, - b_req_idx_cpu=b_req_idx_cpu, - b_ready_cache_len_cpu=b_ready_cache_len_cpu, - b_seq_len_cpu=b_seq_len_cpu, + req_list=req_list, + ready_list=ready_list, + seq_list=seq_list, + mem_indexes=mem_indexes, ) self.prepare_prefill_compress_slots( - b_req_idx, - b_ready_cache_len, - b_seq_len, - b_req_idx_cpu=b_req_idx_cpu, - b_ready_cache_len_cpu=b_ready_cache_len_cpu, - b_seq_len_cpu=b_seq_len_cpu, + req_list=req_list, + ready_list=ready_list, + seq_list=seq_list, + mem_indexes=mem_indexes, ) return def prepare_decode_swa( self, - b_req_idx_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + seq_list: List[int], mem_indexes: torch.Tensor, + prev_mem_indexes: Optional[torch.Tensor] = None, ) -> None: """decode prep: 回收出窗槽并为本步新 token 分配位置对齐的 swa 槽。当前 query 位置 seq_len-1 的窗口是 [seq_len-W, seq_len-1];回收边界额外保留一个 radix 页 @@ -589,8 +561,6 @@ def prepare_decode_swa( if self.sliding_window is not None: retain = self._swa_retain_len() evict_slots = [] - req_list = b_req_idx_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() for req_idx, seq_len in zip(req_list, seq_list): if req_idx == self.HOLD_REQUEST_ID: continue @@ -605,7 +575,20 @@ def prepare_decode_swa( self._swa_evict_marks[req_idx] = evict_end if evict_slots: self.mem_manager.evict_swa(torch.cat(evict_slots)) - self.mem_manager.alloc_swa_decode(b_req_idx_cpu, b_seq_len_cpu, mem_indexes, self.req_to_token_indexs) + if prev_mem_indexes is None: + prev_meta = g_pin_mem_manager.gen_from_list( + key="dsv4_swa_decode_prev", + data=[x for req_idx, seq_len in zip(req_list, seq_list) for x in (req_idx, seq_len - 2)], + dtype=torch.int64, + ).to(self.req_to_token_indexs.device, non_blocking=True) + prev_meta = prev_meta.view(-1, 2) + prev_mem_indexes = self.req_to_token_indexs[prev_meta[:, 0], prev_meta[:, 1]] + self.mem_manager.alloc_swa_decode( + req_list, + seq_list, + mem_indexes, + prev_mem_indexes, + ) return def init_compress_state(self, req_idx: int): @@ -616,138 +599,41 @@ def init_compress_state(self, req_idx: int): return # ------------------------------------------------------------------ compress slot prep (per step) - def _compress_mapping_alloc(self, ratio: int): - assert self.mem_manager is not None, "DeepSeek-V4 mem manager is not bound yet" - if ratio == 4: - raise AssertionError("DeepSeek-V4 c4 uses page-safe allocation") - if ratio == 128: - return self.mem_manager.full_to_c128_indexs, self.mem_manager.alloc_c128 - raise AssertionError(f"invalid DeepSeek-V4 compress ratio {ratio}") - - def _c4_group_end_full_slots(self, req_rows, entries: torch.Tensor) -> torch.Tensor: - """组末 token 的 full 槽位 id (token 位置 = entry*4+3);req_rows 可为标量 req_idx 或行张量。""" - return self.req_to_token_indexs[req_rows, entries * 4 + 3].long() - def _register_c4_slots(self, full_slots: torch.Tensor, slots: torch.Tensor) -> None: """写入 full->c4 槽映射并按页累加存活计数。""" self.mem_manager.full_to_c4_indexs[full_slots] = slots self.mem_manager.count_c4_slots(slots, 1) - def _scatter_c4_prefill_slots_slow(self, req_idx: int, first: int, last: int) -> None: - """Idempotence fallback for overlapped/repeated c4 prep.""" - page = DSV4_C4_PAGE_SIZE - mapping = self.mem_manager.full_to_c4_indexs - for page_base in range((first // page) * page, last, page): - e0 = max(first, page_base) - e1 = min(last, page_base + page) - entries = torch.arange(e0, e1, dtype=torch.long, device="cuda") - full_slots = self._c4_group_end_full_slots(req_idx, entries) - existing = mapping[full_slots] - missing = existing < 0 - if not bool(missing.any()): - continue - - mapped = torch.nonzero(existing >= 0, as_tuple=False) - if mapped.numel() > 0: - j = int(mapped[0].item()) - base = int(existing[j].item()) - ((e0 + j) % page) - elif e0 > page_base: - prev_full = self.req_to_token_indexs[req_idx, e0 * 4 - 1].long() - prev_slot = int(mapping[prev_full].item()) - assert prev_slot >= 0 and prev_slot % page == (e0 - 1) % page - base = prev_slot - ((e0 - 1) % page) - else: - base = int(self.mem_manager.alloc_c4_pages(1)[0].item()) * page - - slots = (base + entries % page).to(torch.int32) - if mapped.numel() > 0: - assert bool((existing[existing >= 0] == slots[existing >= 0]).all()) - self._register_c4_slots(full_slots[missing], slots[missing]) - return - - def _scatter_c4_prefill_slots(self, req_idx: int, first: int, last: int) -> None: - """为 logical c4 entry [first, last) 分配 page-safe c4 槽。 - - 不变式: logical entry e 映射到 physical_page * 64 + e % 64,同一 logical page - 内 entry 共享 physical_page。这是 DeepGEMM paged MQA logits 直接消费 page table 的前提。 - """ - if last <= first: - return - mapping = self.mem_manager.full_to_c4_indexs - full_slots = self._c4_group_end_full_slots(req_idx, torch.arange(first, last, device="cuda")) - need = mapping[full_slots] < 0 - if not bool(need.any()): - return - if not bool(need.all()): - self._scatter_c4_prefill_slots_slow(req_idx, first, last) - return - self._scatter_c4_prefill_slots_fresh(req_idx, first, last) - return - - def _scatter_c4_prefill_slots_fresh(self, req_idx: int, first: int, last: int) -> None: - """Sync-free fast path for entries [first, last) the caller already knows are all fresh: - page-safe alloc with the continuation base read on-GPU. KvCacheAllocator.alloc is CPU-side - (watermark + pinned buffer, no D2H), so this has no syncs.""" - page = DSV4_C4_PAGE_SIZE - mapping = self.mem_manager.full_to_c4_indexs - entries = torch.arange(first, last, dtype=torch.long, device="cuda") - full_slots = self._c4_group_end_full_slots(req_idx, entries) - first_page = first // page - n_pages = (last - 1) // page - first_page + 1 - bases = torch.empty((n_pages,), dtype=torch.long, device="cuda") - base_start = 0 - if first % page != 0: # chunk starts mid-page -> continue the prev chunk's physical page - prev_full = self.req_to_token_indexs[req_idx, first * 4 - 1].long() - bases[0] = mapping[prev_full].long() - ((first - 1) % page) - base_start = 1 - if n_pages - base_start > 0: - bases[base_start:] = ( - self.mem_manager.alloc_c4_pages(n_pages - base_start).cuda(non_blocking=True).long() * page - ) - page_local = torch.div(entries, page, rounding_mode="floor") - first_page - slots = (bases[page_local] + entries % page).to(torch.int32) - self._register_c4_slots(full_slots, slots) - return + def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list, mem_indexes) -> None: + """Batch c4 prefill scatter from the generic preprocess full-slot layout. - def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list) -> None: - """Whole-batch c4 prefill scatter in O(1) GPU ops (independent of request count). The per-req - loop cost O(N) launches + 2-3 D2H syncs each; here every request's group-end entries are - flattened (ragged) and processed in one gather / idempotency-check / page-alloc / scatter / - count. Falls back to the per-req idempotent path on partial/re-run; preserves the page - invariant (logical entry e -> physical_page*64 + e%64, same logical page shares a physical - page). KvCacheAllocator.alloc is CPU-side so one batched alloc has no D2H.""" + Each group's end token is in the current chunk, so its full slot is addressed directly in + mem_indexes. Only a mid-page continuation reads the previous group's old req-table entry. + New full slots guarantee a fresh mapping; no GPU-to-CPU idempotency check is needed.""" page = DSV4_C4_PAGE_SIZE mapping = self.mem_manager.full_to_c4_indexs device = mapping.device - # host plan (cheap int arithmetic, no GPU). dup req in one call would break vectorized - # continuation ordering -> fall back to the safe per-req path. - plan, seen, duplicate_req = [], set(), False - c4_page_need = 0 + plan = [] + mem_offset = 0 for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): - req_idx = int(req_idx) + q_len = seq_len - ready_len if req_idx == self.HOLD_REQUEST_ID: + mem_offset += q_len continue - first, last = int(ready_len) // 4, int(seq_len) // 4 + first, last = ready_len // 4, seq_len // 4 if last <= first: + mem_offset += q_len continue - c4_page_need += (last - 1) // page - first // page + 1 # 上界=区间触及页数, 复用本循环 - duplicate_req |= req_idx in seen - seen.add(req_idx) - plan.append((req_idx, first, last)) + plan.append((req_idx, ready_len, mem_offset, first, last)) + mem_offset += q_len if not plan: return - # 兑现: 在所有分支(dup/fresh/batched)的 alloc_c4_pages 之前统一腾页 - self._realize_c4_pages(c4_page_need) - if duplicate_req: - for req_idx, first, last in plan: - self._scatter_c4_prefill_slots(req_idx, first, last) - return def to_cuda_long(key, data): return g_pin_mem_manager.gen_from_list(key=key, data=data, dtype=torch.int64).to(device, non_blocking=True) - reqs, firsts, lasts = zip(*plan) + reqs, readies, mem_offsets, firsts, lasts = zip(*plan) counts = [last - first for first, last in zip(firsts, lasts)] first_pages = [first // page for first in firsts] page_counts = [((last - 1) // page) - fp + 1 for last, fp in zip(lasts, first_pages)] @@ -756,59 +642,60 @@ def to_cuda_long(key, data): page_offsets.append(total_pages) total_pages += n_pages total_entries = sum(counts) + cont = [(off, req, first) for off, req, first in zip(page_offsets, reqs, firsts) if first % page != 0] + self._realize_c4_pages(total_pages - len(cont)) - # one pinned H2D copy for all per-request metadata (5 cols), then per-entry ragged expansion + # One pinned H2D copy for all per-request metadata, then per-entry ragged expansion. meta = to_cuda_long( "dsv4_c4_prefill_meta", - [x for row in zip(reqs, firsts, first_pages, counts, page_offsets) for x in row], - ).view(-1, 5) - reqs_t, firsts_t, first_pages_t, counts_t, page_offsets_t = meta.unbind(1) + [x for row in zip(readies, mem_offsets, firsts, first_pages, counts, page_offsets) for x in row], + ).view(-1, 6) + readies_t, mem_offsets_t, firsts_t, first_pages_t, counts_t, page_offsets_t = meta.unbind(1) seg = torch.repeat_interleave(torch.arange(len(plan), device=device), counts_t, output_size=total_entries) seg_starts = counts_t.cumsum(0) - counts_t entries = firsts_t[seg] + torch.arange(total_entries, device=device) - seg_starts[seg] - full_slots = self._c4_group_end_full_slots(reqs_t[seg], entries) - - if not bool((mapping[full_slots] < 0).all()): # the single batched idempotency sync - for req_idx, first, last in plan: - self._scatter_c4_prefill_slots(req_idx, first, last) - return + full_offsets = mem_offsets_t[seg] + entries * 4 + 3 - readies_t[seg] + full_slots = mem_indexes.reshape(-1)[full_offsets] # physical base per logical page: fresh pages from one alloc; mid-page continuations read prev - cont = [(off, req, first) for off, req, first in zip(page_offsets, reqs, firsts) if first % page != 0] if not cont: - page_bases = self.mem_manager.alloc_c4_pages(total_pages).to(device, non_blocking=True).long() * page + page_bases = self.mem_manager.alloc_c4_pages(total_pages).to(device, non_blocking=True) * page else: - page_bases = torch.empty(total_pages, dtype=torch.long, device=device) + page_bases = torch.empty(total_pages, dtype=torch.int32, device=device) new_pos = [ pos for off, n_pages, first in zip(page_offsets, page_counts, firsts) - for pos in range(off + int(first % page != 0), off + n_pages) + for pos in range(off + (first % page != 0), off + n_pages) ] if new_pos: new_pos_t = to_cuda_long("dsv4_c4_prefill_new_pos", new_pos) page_bases[new_pos_t] = ( - self.mem_manager.alloc_c4_pages(len(new_pos)).to(device, non_blocking=True).long() * page + self.mem_manager.alloc_c4_pages(len(new_pos)).to(device, non_blocking=True) * page ) cont_t = to_cuda_long("dsv4_c4_prefill_cont", [x for row in cont for x in row]).view(-1, 3) - prev_slot = mapping[self.req_to_token_indexs[cont_t[:, 1], cont_t[:, 2] * 4 - 1].long()].long() - cont_off = (cont_t[:, 2] - 1) % page - assert bool((prev_slot >= 0).all()) and bool(((prev_slot % page) == cont_off).all()) + prev_slot = mapping[self.req_to_token_indexs[cont_t[:, 1], cont_t[:, 2] * 4 - 1]] + cont_off = ((cont_t[:, 2] - 1) % page).to(torch.int32) page_bases[cont_t[:, 0]] = prev_slot - cont_off page_idx = page_offsets_t[seg] + torch.div(entries, page, rounding_mode="floor") - first_pages_t[seg] - slots = (page_bases[page_idx] + entries % page).to(torch.int32) + slots = page_bases[page_idx] + (entries % page).to(torch.int32) self._register_c4_slots(full_slots, slots) return - def _scatter_c4_decode_slots(self, req_list, seq_list, mem_indexes: torch.Tensor) -> None: + def _scatter_c4_decode_slots( + self, + req_list, + seq_list, + mem_indexes: torch.Tensor, + prev_group_end_mem_indexes: Optional[torch.Tensor] = None, + ) -> None: page = DSV4_C4_PAGE_SIZE mapping = self.mem_manager.full_to_c4_indexs - mem_indexes = mem_indexes.cuda().long().reshape(-1) + mem_indexes = mem_indexes.reshape(-1) - cont_rows, cont_prev_pos, cont_offsets = [], [], [] + cont_rows, cont_prev_pos = [], [] new_rows = [] for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): - req_idx, seq_len = int(req_idx), int(seq_len) if req_idx == self.HOLD_REQUEST_ID or seq_len <= 0 or seq_len % 4 != 0: continue entry = seq_len // 4 - 1 @@ -818,38 +705,35 @@ def _scatter_c4_decode_slots(self, req_list, seq_list, mem_indexes: torch.Tensor else: cont_rows.append(i) cont_prev_pos.append(entry * 4 - 1) - cont_offsets.append(offset) if cont_rows: - req_rows = torch.tensor([req_list[i] for i in cont_rows], dtype=torch.long, device="cuda") - prev_pos = torch.tensor(cont_prev_pos, dtype=torch.long, device="cuda") - prev_full = self.req_to_token_indexs[req_rows, prev_pos].long() + if prev_group_end_mem_indexes is None: + prev_meta = g_pin_mem_manager.gen_from_list( + key="dsv4_c4_decode_prev", + data=[x for row in zip([req_list[i] for i in cont_rows], cont_prev_pos) for x in row], + dtype=torch.int64, + ).to(mapping.device, non_blocking=True) + prev_meta = prev_meta.view(-1, 2) + prev_full = self.req_to_token_indexs[prev_meta[:, 0], prev_meta[:, 1]] + else: + prev_full = prev_group_end_mem_indexes.reshape(-1)[cont_rows] prev_slots = mapping[prev_full] - offsets = torch.tensor(cont_offsets, dtype=torch.int32, device="cuda") - assert bool((prev_slots >= 0).all()) - assert bool(((prev_slots % page) == (offsets - 1)).all()) - self._register_c4_slots(mem_indexes[cont_rows], (prev_slots + 1).to(torch.int32)) + self._register_c4_slots(mem_indexes[cont_rows], prev_slots + 1) if new_rows: self._realize_c4_pages(len(new_rows)) # 兑现: 精确需求, 复用已算的 new_rows - pages = self.mem_manager.alloc_c4_pages(len(new_rows)).cuda(non_blocking=True).long() - self._register_c4_slots(mem_indexes[new_rows], (pages * page).to(torch.int32)) + pages = self.mem_manager.alloc_c4_pages(len(new_rows)).to(mapping.device, non_blocking=True) + self._register_c4_slots(mem_indexes[new_rows], pages * page) return - def _scatter_compress_slots(self, ratio: int, full_slots: torch.Tensor) -> None: - """为组末 full 槽位分配压缩槽并写入映射。已映射(>=0)的行跳过——重复 prep 幂等。""" + def _scatter_c128_slots(self, full_slots: torch.Tensor) -> None: + """为本批新组末 full 槽分配 c128 槽并写入映射。""" if full_slots.numel() == 0: return - mapping, alloc = self._compress_mapping_alloc(ratio) - full_slots = full_slots.cuda().long().reshape(-1) - # 去重: 同批重复键会让后写覆盖先写,先分配的压缩槽成为孤儿(allocator 泄漏)。 - need = torch.unique(full_slots[mapping[full_slots] < 0]) - if need.numel() == 0: - return - if ratio == 128: # _scatter_compress_slots 仅用于 c128; 兑现其槽(复用已算的 need.numel()) - self._realize_c128_slots(int(need.numel())) - new_slots = alloc(need.numel()).cuda(non_blocking=True).to(torch.int32) - mapping[need] = new_slots + full_slots = full_slots.reshape(-1) + self._realize_c128_slots(full_slots.numel()) + new_slots = self.mem_manager.alloc_c128(full_slots.numel()).cuda(non_blocking=True) + self.mem_manager.full_to_c128_indexs[full_slots] = new_slots return def _realize_c4_pages(self, need_pages: int) -> None: @@ -877,54 +761,59 @@ def _realize_c128_slots(self, need_slots: int) -> None: def prepare_prefill_compress_slots( self, - b_req_idx: torch.Tensor, - b_ready_cache_len: torch.Tensor, - b_seq_len: torch.Tensor, - b_req_idx_cpu: torch.Tensor, - b_ready_cache_len_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + ready_list: List[int], + seq_list: List[int], + mem_indexes: torch.Tensor, ) -> None: """prefill prep: 为本 chunk 内的组末 token(位置 (g+1)*ratio-1 ∈ [ready, seq))分配压缩槽, - scatter 进 full_to_c4/c128_indexs。必须在 init_req_to_token_indexes 之后(组末 full 槽 - 从 req_to_token_indexs 取)、attention metadata 构建之前调用。""" + 组末 full 槽直接从 generic preprocess 的 mem_indexes 取。""" if self.n_c4 == 0 and self.n_c128 == 0: return - req_list = b_req_idx_cpu.tolist() - ready_list = b_ready_cache_len_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() if self.n_c4 > 0: - self._scatter_c4_prefill_slots_batched(req_list, ready_list, seq_list) + self._scatter_c4_prefill_slots_batched(req_list, ready_list, seq_list, mem_indexes) if self.n_c128 > 0: ratio = 128 - end_slots = [] + full_offsets = [] + mem_offset = 0 for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): - req_idx = int(req_idx) + q_len = seq_len - ready_len if req_idx == self.HOLD_REQUEST_ID: + mem_offset += q_len continue - first, last = int(ready_len) // ratio, int(seq_len) // ratio + first, last = ready_len // ratio, seq_len // ratio if last > first: - ends = self.req_to_token_indexs[req_idx, ratio - 1 : last * ratio : ratio] - end_slots.append(ends[first:]) - if end_slots: - self._scatter_compress_slots(ratio, torch.cat(end_slots)) + full_offsets.extend( + mem_offset + (entry + 1) * ratio - 1 - ready_len for entry in range(first, last) + ) + mem_offset += q_len + if full_offsets: + offsets = g_pin_mem_manager.gen_from_list( + key="dsv4_c128_prefill_offsets", data=full_offsets, dtype=torch.int64 + ).to(mem_indexes.device, non_blocking=True) + self._scatter_c128_slots(mem_indexes.reshape(-1)[offsets]) return def prepare_decode_compress_slots( self, - b_req_idx_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, + req_list: List[int], + seq_list: List[int], mem_indexes: torch.Tensor, + prev_group_end_mem_indexes: Optional[torch.Tensor] = None, ) -> None: """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 组末 full 槽即本步的 mem_index。 从 CPU 镜像读 seq_len/req_idx(host 算术,无 D2H);非关组步 rows 为空 => 不调 _scatter,零同步。""" if self.n_c4 == 0 and self.n_c128 == 0: return - req_list = b_req_idx_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() if self.n_c4 > 0: - self._scatter_c4_decode_slots(req_list, seq_list, mem_indexes) + self._scatter_c4_decode_slots( + req_list, + seq_list, + mem_indexes, + prev_group_end_mem_indexes=prev_group_end_mem_indexes, + ) if self.n_c128 > 0: ratio = 128 @@ -934,7 +823,7 @@ def prepare_decode_compress_slots( if req_idx != self.HOLD_REQUEST_ID and seq_len > 0 and seq_len % ratio == 0 ] if rows: - self._scatter_compress_slots(ratio, mem_indexes.reshape(-1)[rows]) + self._scatter_c128_slots(mem_indexes.reshape(-1)[rows]) return def alloc(self): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 50e2e22de1..ab8be8c613 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -6,6 +6,7 @@ import torch from lightllm.models.registry import ModelRegistry from lightllm.models.llama.model import LlamaTpPartModel +from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import ( @@ -127,6 +128,53 @@ def _init_custom(self): ) return + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: + """Commit DSV4 derived slots before BaseModel pads or scatters the generic input.""" + if model_input.is_prefill and self.is_mtp_draft_model: + return + if model_input.mem_indexes_cpu is None: + return + if model_input.mem_indexes is None: + model_input.mem_indexes = model_input.mem_indexes_cpu.cuda(non_blocking=True) + + if model_input.is_prefill: + self.req_manager.prepare_prefill( + b_req_idx_cpu=model_input.b_req_idx_cpu, + b_ready_cache_len_cpu=model_input.b_ready_cache_len, + b_seq_len_cpu=model_input.b_seq_len_cpu, + mem_indexes=model_input.mem_indexes, + ) + return + + if model_input.mtp_decode_slot_prepare_indices == (): + return + self.req_manager.prepare_decode( + model_input.b_req_idx_cpu, + model_input.b_seq_len_cpu, + model_input.b_mtp_index_cpu, + model_input.mem_indexes, + model_input.mtp_decode_slot_prepare_indices, + prepare_compress_slots=not self.is_mtp_draft_model, + ) + return + + @torch.no_grad() + def forward(self, model_input: ModelInput): + self._prepare_dsv4_slots(model_input) + return super().forward(model_input) + + @torch.no_grad() + def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: ModelInput): + self._prepare_dsv4_slots(model_input0) + self._prepare_dsv4_slots(model_input1) + return super().microbatch_overlap_prefill(model_input0, model_input1) + + @torch.no_grad() + def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: ModelInput): + self._prepare_dsv4_slots(model_input0) + self._prepare_dsv4_slots(model_input1) + return super().microbatch_overlap_decode(model_input0, model_input1) + def _init_to_get_rotary(self): # Interleaved (GPT-J) rope. Build complex64 freqs_cis tables (_freqs_cis_*) following the # gemma4 two-variant convention; the fused sglang q kernel consumes them directly, while diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index bf87c845c2..40d593d0e8 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -27,8 +27,6 @@ has_vision_module, is_linear_att_mixed_model, auto_set_max_req_total_len, - get_model_type, - get_config_json, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args From 388dd29de8a000500279ad273c06a6d7874f7cf1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 08:06:06 +0000 Subject: [PATCH 075/214] autotune --- ...=64,dtype=torch.bfloat16}_NVIDIA_H200.json | 74 +++ ...out_dtype=torch.bfloat16}_NVIDIA_H200.json | 554 ++++++++++++++++++ 2 files changed, 628 insertions(+) create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json new file mode 100644 index 0000000000..e40e19975e --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 4, + "num_warps": 8 + }, + "100": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 3, + "num_warps": 1 + }, + "1024": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 4, + "num_stages": 1, + "num_warps": 2 + }, + "128": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 3, + "num_warps": 1 + }, + "16": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 4 + }, + "2048": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 2, + "num_stages": 5, + "num_warps": 1 + }, + "256": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 5, + "num_warps": 2 + }, + "32": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 5, + "num_warps": 4 + }, + "4096": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 3, + "num_warps": 1 + }, + "64": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 1, + "num_warps": 1 + }, + "8": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 1 + }, + "8192": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 1, + "num_stages": 2, + "num_warps": 2 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json new file mode 100644 index 0000000000..d0ea86fac3 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json @@ -0,0 +1,554 @@ +{ + "1": { + "BLOCK_M": 256, + "BLOCK_N": 64, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "100": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "1024": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "1152": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "12160": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1280": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "13056": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13312": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13696": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13824": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13952": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1408": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14080": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14336": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14464": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14720": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14848": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14976": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "15104": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1536": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "16": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "1664": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1920": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2048": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2176": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2304": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "24192": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2432": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "24960": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25088": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25344": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "256": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "2560": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25600": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25856": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25984": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26112": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26368": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26624": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2688": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27136": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27520": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27776": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2816": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "28160": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "28800": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2944": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3072": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "32": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "3200": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3328": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3456": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3584": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3712": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "384": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3840": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3968": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4096": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4224": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4352": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4480": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "46336": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "46592": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "48896": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "49152": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "49280": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50560": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50688": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50944": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "512": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "51328": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "52608": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "53248": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "53632": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54400": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54656": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "55040": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "640": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "7296": { + "BLOCK_M": 32, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7552": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "768": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7808": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7936": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8": { + "BLOCK_M": 1, + "BLOCK_N": 32, + "NUM_STAGES": 4, + "num_warps": 8 + }, + "8064": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8192": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8320": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8448": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8704": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "896": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + } +} \ No newline at end of file From 6bfd98f4074034a8f26c33c548ce6f4bb976452a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 13 Jul 2026 08:10:36 +0000 Subject: [PATCH 076/214] delete third_party --- .../deepseek4_mem_manager.py | 7 +- .../layer_infer/transformer_layer_infer.py | 27 +- lightllm/models/deepseek_v4/model.py | 6 +- .../triton_kernel/csrc/norm_rope.cu | 622 +++++++++++++ .../triton_kernel/csrc/topk_transform.cu | 345 +++++++ .../triton_kernel/norm_rope_cuda.py | 98 ++ .../triton_kernel/topk_transform.py | 61 ++ lightllm/third_party/__init__.py | 1 - lightllm/third_party/sglang_jit/LICENSE | 201 ---- lightllm/third_party/sglang_jit/README.md | 13 - lightllm/third_party/sglang_jit/__init__.py | 1 - .../sglang_jit/csrc/deepseek_v4/c128.cuh | 522 ----------- .../csrc/deepseek_v4/c128_online.cuh | 726 --------------- .../csrc/deepseek_v4/c128_online_v2.cuh | 875 ------------------ .../sglang_jit/csrc/deepseek_v4/c128_v2.cuh | 448 --------- .../sglang_jit/csrc/deepseek_v4/c4.cuh | 549 ----------- .../sglang_jit/csrc/deepseek_v4/c4_v2.cuh | 405 -------- .../sglang_jit/csrc/deepseek_v4/c_plan.cuh | 839 ----------------- .../sglang_jit/csrc/deepseek_v4/common.cuh | 208 ----- .../csrc/deepseek_v4/fused_norm_rope.cuh | 254 ----- .../csrc/deepseek_v4/fused_norm_rope_v2.cuh | 643 ------------- .../sglang_jit/csrc/deepseek_v4/hash_topk.cuh | 214 ----- .../csrc/deepseek_v4/hisparse_transfer.cuh | 82 -- .../csrc/deepseek_v4/main_norm_rope.cuh | 845 ----------------- .../deepseek_v4/mega_moe_pre_dispatch.cuh | 219 ----- .../csrc/deepseek_v4/paged_mqa_metadata.cuh | 119 --- .../sglang_jit/csrc/deepseek_v4/rope.cuh | 169 ---- .../silu_and_mul_masked_post_quant.cuh | 540 ----------- .../sglang_jit/csrc/deepseek_v4/store.cuh | 205 ---- .../sglang_jit/csrc/deepseek_v4/topk_v1.cuh | 340 ------- .../sglang_jit/csrc/deepseek_v4/topk_v2.cuh | 493 ---------- .../third_party/sglang_jit/dsv4/__init__.py | 8 - .../sglang_jit/dsv4/elementwise.py | 215 ----- lightllm/third_party/sglang_jit/dsv4/topk.py | 92 -- lightllm/third_party/sglang_jit/dsv4/utils.py | 2 - .../sglang_jit/include/sgl_kernel/atomic.cuh | 35 - .../sglang_jit/include/sgl_kernel/cta.cuh | 40 - .../sgl_kernel/deepseek_v4/compress.cuh | 37 - .../sgl_kernel/deepseek_v4/compress_v2.cuh | 99 -- .../sgl_kernel/deepseek_v4/fp8_utils.cuh | 112 --- .../sgl_kernel/deepseek_v4/kvcacheio.cuh | 96 -- .../sgl_kernel/deepseek_v4/topk/cluster.cuh | 257 ----- .../sgl_kernel/deepseek_v4/topk/common.cuh | 176 ---- .../sgl_kernel/deepseek_v4/topk/ptx.cuh | 54 -- .../sgl_kernel/deepseek_v4/topk/register.cuh | 302 ------ .../sgl_kernel/deepseek_v4/topk/streaming.cuh | 213 ----- .../include/sgl_kernel/distributed/common.cuh | 120 --- .../distributed/custom_all_reduce.cuh | 354 ------- .../sglang_jit/include/sgl_kernel/ffi.h | 104 --- .../include/sgl_kernel/impl/norm.cuh | 168 ---- .../sglang_jit/include/sgl_kernel/math.cuh | 71 -- .../sglang_jit/include/sgl_kernel/runtime.cuh | 86 -- .../include/sgl_kernel/scalar_type.hpp | 334 ------- .../include/sgl_kernel/source_location.h | 40 - .../sglang_jit/include/sgl_kernel/tensor.h | 605 ------------ .../sglang_jit/include/sgl_kernel/tile.cuh | 62 -- .../sglang_jit/include/sgl_kernel/type.cuh | 120 --- .../sglang_jit/include/sgl_kernel/utils.cuh | 333 ------- .../sglang_jit/include/sgl_kernel/utils.h | 186 ---- .../sglang_jit/include/sgl_kernel/vec.cuh | 118 --- .../sglang_jit/include/sgl_kernel/warp.cuh | 56 -- lightllm/third_party/sglang_jit/jit_utils.py | 432 --------- .../third_party/sglang_jit/runtime_utils.py | 5 - 63 files changed, 1147 insertions(+), 13862 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu create mode 100644 lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu create mode 100644 lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/topk_transform.py delete mode 100644 lightllm/third_party/__init__.py delete mode 100755 lightllm/third_party/sglang_jit/LICENSE delete mode 100644 lightllm/third_party/sglang_jit/README.md delete mode 100644 lightllm/third_party/sglang_jit/__init__.py delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh delete mode 100644 lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/dsv4/__init__.py delete mode 100644 lightllm/third_party/sglang_jit/dsv4/elementwise.py delete mode 100644 lightllm/third_party/sglang_jit/dsv4/topk.py delete mode 100644 lightllm/third_party/sglang_jit/dsv4/utils.py delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh delete mode 100644 lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh delete mode 100644 lightllm/third_party/sglang_jit/jit_utils.py delete mode 100644 lightllm/third_party/sglang_jit/runtime_utils.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 1e468a4eef..883bb59937 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -684,11 +684,12 @@ def pack_mla_kv_to_cache_fused_norm_rope( positions: torch.Tensor, ): """同 pack_mla_kv_to_cache,但 rmsnorm + 尾部交错 rope 融合进写入 kernel - (sglang fused_k_norm_rope_flashmla,即 sglang _compute_kv_to_cache 的池侧), - 省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" + 并省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" if kv.shape[0] == 0: return - from lightllm.third_party.sglang_jit.dsv4 import fused_k_norm_rope_flashmla + from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import ( + fused_k_norm_rope_flashmla, + ) swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] swa_slots = torch.where(swa_slots < 0, torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), swa_slots) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 3cd3ec7010..3e0315bcb5 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -16,7 +16,7 @@ from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo import deep_gemm -from lightllm.third_party.sglang_jit.dsv4 import topk_transform_512 +from lightllm.models.deepseek_v4.triton_kernel.topk_transform import topk_transform_512 _C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 @@ -151,21 +151,21 @@ def _get_qkv( infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, ): - from lightllm.third_party.sglang_jit.dsv4 import fused_q_norm_rope + from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import fused_q_norm_rope input = self._tpsp_allgather(input=input, infer_state=infer_state) T = input.shape[0] # wq_a and wkv share `input` -> one fused fp8 GEMM, split [q_lora_rank | head_dim]. qa is a - # row-strided view (rmsnorm honors stride(0)); kv feeds a sglang jit kernel -> contiguous. + # row-strided view (rmsnorm honors stride(0)); kv feeds the fused cache writer -> contiguous. qkv = layer_weight.wq_a_wkv_.mm(input) qa = layer_weight.q_norm_(qkv[:, : -self.head_dim_], eps=self.eps_) q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, - # fused in one sglang dsv4 jit kernel (fp32 norm/rotation, bf16 in between -- same as eager). + # fused in one DSV4 CUDA kernel (fp32 norm/rotation, bf16 in between -- same as eager). q = self.alloc_tensor(q_in.shape, dtype=q_in.dtype, device=q_in.device) fused_q_norm_rope(q_in, q, self.eps_, self.freqs_cis, infer_state.position_ids) - # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 sglang jit kernel 完成 - # (同 sglang _compute_kv_to_cache),替代 eager norm/rope/cat + _post_cache_kv。 + # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 DSV4 CUDA kernel 完成, + # 替代 eager norm/rope/cat + _post_cache_kv。 # bf16 kv 中间量没有其他消费者: flashmla 路径注意力读 cache,压缩器/indexer 取 x。 infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, @@ -394,11 +394,6 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV def _select_experts( self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - return self._select_experts_vllm(logits, infer_state, layer_weight) - - def _select_experts_vllm( - self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): M = logits.shape[0] bias = None @@ -624,7 +619,9 @@ def _indexer_q_weight( # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: # rotary_emb_fwd + hadamard_transform + act_quant + weights mul). freqs_cis is the compress rope # table (same one the main compress-layer Q path uses); positions indexed inside the kernel. - from lightllm.third_party.sglang_jit.dsv4.elementwise import fused_q_indexer_rope_hadamard_quant + from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import ( + fused_q_indexer_rope_hadamard_quant, + ) token_num = q_lora.shape[0] if x.shape[0] != token_num: @@ -638,7 +635,11 @@ def _indexer_q_weight( token_num, self.index_n_heads ) # [T, H] raw idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( - idx_q, raw_w, self.indexer_weight_scale, self.freqs_cis, infer_state.position_ids + idx_q, + raw_w, + self.indexer_weight_scale, + self.freqs_cis, + infer_state.position_ids, ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded return idx_q_fp8, weights.squeeze(-1).contiguous() diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index ab8be8c613..567b14ea49 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -177,9 +177,9 @@ def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: Mode def _init_to_get_rotary(self): # Interleaved (GPT-J) rope. Build complex64 freqs_cis tables (_freqs_cis_*) following the - # gemma4 two-variant convention; the fused sglang q kernel consumes them directly, while - # _cos_cached_*/_sin_cached_* are .real/.imag views of the same storage for the kv rope, - # inverse rope and compressor paths (deepseek2's interleaved triton rotary_emb_fwd). + # gemma4 two-variant convention; the fused CUDA Q/K kernels consume them directly, while + # _cos_cached_*/_sin_cached_* are .real/.imag views of the same storage for the inverse + # rope and compressor paths (deepseek2's interleaved triton rotary_emb_fwd). # Sliding-window and compressed layers both use DeepSeek YaRN correction; only the # RoPE base differs (rope_theta vs compress_rope_theta), matching SGLang/vLLM. # Kept fp32 for accuracy (the apply upcasts anyway). diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu new file mode 100644 index 0000000000..b0a2963a53 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu @@ -0,0 +1,622 @@ +// Copyright 2023-2024 SGLang Team +// SPDX-License-Identifier: Apache-2.0 +// +// DeepSeek-V4 main Q/K and indexer-Q fused kernels. +// +// Adapted from SGLang commit 8cea0473ea5299bc04885f8f6ba71269415a39b5, +// python/sglang/jit_kernel/csrc/deepseek_v4/main_norm_rope.cuh. This local +// port keeps the device math, warp mapping, launch bounds, and Hopper PDL +// protocol, while replacing tvm::ffi and the generic SGLang headers with a +// small torch cpp_extension binding specialized for LightLLM's DSV4 shapes. + +#include +#include +#include +#include + +#include +#include +#include + +#include +#include +#include + +namespace { + +constexpr uint32_t kWarpThreads = 32; +constexpr uint32_t kFusedQBlockSize = 128; +constexpr uint32_t kFusedQNumWarps = kFusedQBlockSize / kWarpThreads; +constexpr uint32_t kFusedKBlockSize = 256; +constexpr uint32_t kFusedKNumWarps = kFusedKBlockSize / kWarpThreads; + +constexpr int64_t kMainHeadDim = 512; +constexpr int64_t kMainRopeDim = 64; +constexpr int64_t kMainNopeDim = kMainHeadDim - kMainRopeDim; +constexpr int64_t kIndexerHeadDim = 128; +constexpr int64_t kIndexerRopeDim = 64; +constexpr uint32_t kFlashMLAPageSize = 128; +constexpr int32_t kFlashMLAPageBits = 7; +constexpr int64_t kFlashMLADataBytes = 576; +constexpr int64_t kFlashMLAScaleBytes = 8; +constexpr int64_t kFlashMLABytesPerToken = kFlashMLADataBytes + kFlashMLAScaleBytes; +constexpr int64_t kFlashMLAPageBytes = + ((kFlashMLABytesPerToken * kFlashMLAPageSize + kFlashMLADataBytes - 1) / kFlashMLADataBytes) * + kFlashMLADataBytes; + +template +struct alignas(sizeof(T) * N) AlignedVector { + T data[N]; + + __device__ T& operator[](int i) { return data[i]; } + __device__ const T& operator[](int i) const { return data[i]; } +}; + +template +__device__ __forceinline__ T warp_reduce_sum(T value) { +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + value += __shfl_xor_sync(0xffffffffu, value, mask, 32); + } + return value; +} + +template +__device__ __forceinline__ T warp_reduce_sum_width(T value) { + static_assert(NumThreads > 0 && NumThreads <= 32 && (NumThreads & (NumThreads - 1)) == 0); +#pragma unroll + for (int mask = NumThreads / 2; mask > 0; mask >>= 1) { + value += __shfl_xor_sync(0xffffffffu, value, mask, 32); + } + return value; +} + +template +__device__ __forceinline__ T warp_reduce_max(T value) { +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + value = fmaxf(value, __shfl_xor_sync(0xffffffffu, value, mask, 32)); + } + return value; +} + +template +__device__ __forceinline__ void pdl_wait_primary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) { + asm volatile("griddepcontrol.wait;" ::: "memory"); + } +#endif +} + +template +__device__ __forceinline__ void pdl_trigger_secondary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) { + asm volatile("griddepcontrol.launch_dependents;" :::); + } +#endif +} + +__device__ __forceinline__ int32_t cast_to_ue8m0(float x) { + const uint32_t u = __float_as_uint(x); + const int32_t exp = static_cast((u >> 23) & 0xffu); + const uint32_t mantissa = u & 0x7fffffu; + return exp + (mantissa != 0); +} + +__device__ __forceinline__ float inv_scale_ue8m0(int32_t exp) { + return __uint_as_float(static_cast(127 + 127 - exp) << 23); +} + +__device__ __forceinline__ uint16_t pack_fp8(float x, float y) { + constexpr float kFp8Max = 448.0f; + const float2 values = { + fmaxf(fminf(x, kFp8Max), -kFp8Max), + fmaxf(fminf(y, kFp8Max), -kFp8Max), + }; + return __nv_cvt_float2_to_fp8x2(values, __NV_SATFINITE, __NV_E4M3); +} + +template +void launch_kernel(Kernel kernel, dim3 grid, dim3 block, cudaStream_t stream, bool enable_pdl, Args... args) { + cudaLaunchConfig_t config{}; + config.gridDim = grid; + config.blockDim = block; + config.dynamicSmemBytes = 0; + config.stream = stream; + + cudaLaunchAttribute attribute{}; + if (enable_pdl) { + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = true; + config.attrs = &attribute; + config.numAttrs = 1; + } + + C10_CUDA_CHECK(cudaLaunchKernelEx(&config, kernel, args...)); +} + +bool device_supports_pdl() { + return at::cuda::getCurrentDeviceProperties()->major >= 9; +} + +void check_cuda_tensor(const at::Tensor& tensor, const char* name) { + TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor"); +} + +void check_same_device(const at::Tensor& reference, const at::Tensor& tensor, const char* name) { + check_cuda_tensor(tensor, name); + TORCH_CHECK(tensor.get_device() == reference.get_device(), name, " must be on the same CUDA device as the input"); +} + +struct FusedQNormRopeParams { + const __nv_bfloat16* q_input; + __nv_bfloat16* q_output; + const float* freqs_cis; + const void* positions; + int64_t q_input_stride_batch; + int64_t q_output_stride_batch; + uint32_t batch_size; + uint32_t num_q_heads; + float eps; +}; + +template +__global__ __launch_bounds__(kFusedQBlockSize, 16) void fused_q_norm_rope_kernel( + const __grid_constant__ FusedQNormRopeParams params) { + constexpr int64_t kVecSize = 8; + constexpr int64_t kLocalSize = kMainHeadDim / (kWarpThreads * kVecSize); + constexpr uint32_t kRopeVecs = kMainRopeDim / kVecSize; + using Storage = AlignedVector<__nv_bfloat16, kVecSize>; + + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + const uint32_t work_id = blockIdx.x * kFusedQNumWarps + warp_id; + const uint32_t total_works = params.batch_size * params.num_q_heads; + if (work_id >= total_works) return; + + const uint32_t batch_id = work_id / params.num_q_heads; + const uint32_t head_id = work_id % params.num_q_heads; + const auto* input_ptr = + params.q_input + batch_id * params.q_input_stride_batch + head_id * kMainHeadDim; + auto* output_ptr = + params.q_output + batch_id * params.q_output_stride_batch + head_id * kMainHeadDim; + const int32_t position = static_cast(static_cast(params.positions)[batch_id]); + + __shared__ Storage rope_storage[kFusedQNumWarps][kRopeVecs]; + + pdl_wait_primary(); + + Storage input_vec[kLocalSize]; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + input_vec[i] = reinterpret_cast(input_ptr)[i * kWarpThreads + lane_id]; + } + + const float2 freq = reinterpret_cast(params.freqs_cis + position * kMainRopeDim)[lane_id]; + + float sum_of_squares = 0.0f; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + const float x = __bfloat162float(input_vec[i][j]); + sum_of_squares += x * x; + } + } + sum_of_squares = warp_reduce_sum(sum_of_squares); + const float norm_factor = rsqrtf(sum_of_squares / static_cast(kMainHeadDim) + params.eps); + +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { +#pragma unroll + for (int j = 0; j < kVecSize; ++j) { + input_vec[i][j] = __float2bfloat16_rn(__bfloat162float(input_vec[i][j]) * norm_factor); + } + } + + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeVecs; +#pragma unroll + for (int i = 0; i < kLocalSize; ++i) { + if (i == kLocalSize - 1 && is_rope_lane) { + rope_storage[warp_id][lane_id - (kWarpThreads - kRopeVecs)] = input_vec[i]; + } else { + reinterpret_cast(output_ptr)[i * kWarpThreads + lane_id] = input_vec[i]; + } + } + __syncwarp(); + + pdl_trigger_secondary(); + + const auto elem = reinterpret_cast(rope_storage[warp_id])[lane_id]; + const float2 x = __bfloat1622float2(elem); + const float out_real = x.x * freq.x - x.y * freq.y; + const float out_imag = x.x * freq.y + x.y * freq.x; + reinterpret_cast<__nv_bfloat162*>(output_ptr + kMainNopeDim)[lane_id] = + __floats2bfloat162_rn(out_real, out_imag); +} + +struct FusedKNormRopeFlashMLAParams { + const __nv_bfloat16* kv; + const __nv_bfloat16* kv_weight; + const float* freqs_cis; + const void* positions; + const int32_t* out_loc; + uint8_t* kvcache; + int64_t kv_stride_batch; + uint32_t batch_size; + float eps; +}; + +template +__global__ __launch_bounds__(kFusedKBlockSize, 8) void fused_k_norm_rope_flashmla_kernel( + const __grid_constant__ FusedKNormRopeFlashMLAParams params) { + using Storage = AlignedVector<__nv_bfloat16, 2>; + + const uint32_t tx = threadIdx.x; + const uint32_t warp_id = tx / kWarpThreads; + const uint32_t lane_id = tx % kWarpThreads; + const uint32_t work_id = blockIdx.x; + if (work_id >= params.batch_size) return; + + const auto* input_ptr = params.kv + work_id * params.kv_stride_batch; + const int32_t position = static_cast(static_cast(params.positions)[work_id]); + const int32_t out_loc = params.out_loc[work_id]; + const float* freqs_cis = params.freqs_cis + position * kMainRopeDim; + + pdl_wait_primary(); + + const Storage input_vec = reinterpret_cast(input_ptr)[tx]; + const Storage weight_vec = reinterpret_cast(params.kv_weight)[tx]; + float2 data; + float2 freq{}; + if (warp_id == kFusedKNumWarps - 1) { + freq = reinterpret_cast(freqs_cis)[lane_id]; + } + + float sum_of_squares = 0.0f; + const float input_x = __bfloat162float(input_vec[0]); + const float input_y = __bfloat162float(input_vec[1]); + sum_of_squares += input_x * input_x; + sum_of_squares += input_y * input_y; + const float warp_sum = warp_reduce_sum(sum_of_squares); + + __shared__ float partial_sums[kFusedKNumWarps]; + if (lane_id == 0) partial_sums[warp_id] = warp_sum; + __syncthreads(); + sum_of_squares = warp_reduce_sum_width(partial_sums[lane_id % kFusedKNumWarps]); + const float norm_factor = rsqrtf(sum_of_squares / static_cast(kMainHeadDim) + params.eps); + data.x = input_x * norm_factor * __bfloat162float(weight_vec[0]); + data.y = input_y * norm_factor * __bfloat162float(weight_vec[1]); + + const int32_t page = out_loc >> kFlashMLAPageBits; + const int32_t offset = out_loc & (kFlashMLAPageSize - 1); + auto* page_ptr = params.kvcache + static_cast(page) * kFlashMLAPageBytes; + auto* value_ptr = page_ptr + static_cast(offset) * kFlashMLADataBytes; + + pdl_trigger_secondary(); + + if (warp_id == kFusedKNumWarps - 1) { + const float out_real = data.x * freq.x - data.y * freq.y; + const float out_imag = data.x * freq.y + data.y * freq.x; + reinterpret_cast<__nv_bfloat162*>(value_ptr + kMainNopeDim)[lane_id] = + __floats2bfloat162_rn(out_real, out_imag); + } else { + const float abs_max = warp_reduce_max(fmaxf(fabsf(data.x), fabsf(data.y))); + const float scale_raw = fmaxf(1.0e-4f, abs_max) / 448.0f; + const int32_t scale_ue8m0 = cast_to_ue8m0(scale_raw); + const float inv_scale = inv_scale_ue8m0(scale_ue8m0); + reinterpret_cast(value_ptr)[tx] = pack_fp8(data.x * inv_scale, data.y * inv_scale); + if (lane_id == 0) { + auto* scale_ptr = page_ptr + kFlashMLAPageSize * kFlashMLADataBytes + offset * kFlashMLAScaleBytes; + scale_ptr[warp_id] = static_cast(scale_ue8m0); + } + } +} + +struct FusedQIndexerParams { + const __nv_bfloat16* q_input; + uint8_t* q_fp8; + const __nv_bfloat16* weight; + float* weights_out; + float weight_scale; + const float* freqs_cis; + const void* positions; + uint32_t batch_size; + uint32_t num_heads; +}; + +template +__global__ __launch_bounds__(kFusedQBlockSize, 16) void fused_q_indexer_rope_hadamard_quant_kernel( + const __grid_constant__ FusedQIndexerParams params) { + constexpr int64_t kVecSize = 4; + constexpr uint32_t kRopeVecs = kIndexerRopeDim / kVecSize; + using Storage = AlignedVector<__nv_bfloat16, kVecSize>; + using Float4 = AlignedVector; + + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t lane_id = threadIdx.x % kWarpThreads; + const uint32_t work_id = blockIdx.x * kFusedQNumWarps + warp_id; + const uint32_t total_works = params.batch_size * params.num_heads; + if (work_id >= total_works) return; + + const uint32_t batch_id = work_id / params.num_heads; + const int32_t position = static_cast(static_cast(params.positions)[batch_id]); + const auto* input_ptr = params.q_input + static_cast(work_id) * kIndexerHeadDim; + const float* freqs_cis = params.freqs_cis + position * kIndexerRopeDim; + const bool is_rope_lane = lane_id >= kWarpThreads - kRopeVecs; + + pdl_wait_primary(); + + const float weight_value = __bfloat162float(params.weight[work_id]); + const Storage input_vec = reinterpret_cast(input_ptr)[lane_id]; + Float4 data; + Float4 freq{}; +#pragma unroll + for (int i = 0; i < kVecSize; ++i) data[i] = __bfloat162float(input_vec[i]); + if (is_rope_lane) { + freq = reinterpret_cast(freqs_cis)[lane_id - (kWarpThreads - kRopeVecs)]; + const float x_real = data[0]; + const float x_imag = data[1]; + const float y_real = data[2]; + const float y_imag = data[3]; + data[0] = x_real * freq[0] - x_imag * freq[1]; + data[1] = x_real * freq[1] + x_imag * freq[0]; + data[2] = y_real * freq[2] - y_imag * freq[3]; + data[3] = y_real * freq[3] + y_imag * freq[2]; + } + + pdl_trigger_secondary(); + + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a1; + data[1] = a0 - a1; + data[2] = a2 + a3; + data[3] = a2 - a3; + } + { + const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; + data[0] = a0 + a2; + data[1] = a1 + a3; + data[2] = a0 - a2; + data[3] = a1 - a3; + } +#pragma unroll + for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { +#pragma unroll + for (int i = 0; i < kVecSize; ++i) { + const float other = __shfl_xor_sync(0xffffffffu, data[i], mask, kWarpThreads); + data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); + } + } + constexpr float kHadamardScale = 0.08838834764831845f; // 1 / sqrt(128) +#pragma unroll + for (int i = 0; i < kVecSize; ++i) data[i] *= kHadamardScale; + + float local_max = fabsf(data[0]); +#pragma unroll + for (int i = 1; i < kVecSize; ++i) local_max = fmaxf(local_max, fabsf(data[i])); + const float abs_max = warp_reduce_max(local_max); + const float scale = fmaxf(1.0e-4f, abs_max) / 448.0f; + const float inv_scale = 1.0f / scale; + AlignedVector result; + result[0] = pack_fp8(data[0] * inv_scale, data[1] * inv_scale); + result[1] = pack_fp8(data[2] * inv_scale, data[3] * inv_scale); + auto* output = reinterpret_cast*>( + params.q_fp8 + static_cast(work_id) * kIndexerHeadDim); + output[lane_id] = result; + params.weights_out[work_id] = weight_value * params.weight_scale * scale; +} + +template +void launch_q_norm(const FusedQNormRopeParams& params, cudaStream_t stream) { + const uint32_t works = params.batch_size * params.num_q_heads; + const dim3 grid((works + kFusedQNumWarps - 1) / kFusedQNumWarps); + launch_kernel(fused_q_norm_rope_kernel, grid, kFusedQBlockSize, stream, UsePDL, params); +} + +template +void launch_k_norm(const FusedKNormRopeFlashMLAParams& params, cudaStream_t stream) { + launch_kernel( + fused_k_norm_rope_flashmla_kernel, params.batch_size, kFusedKBlockSize, stream, UsePDL, params); +} + +template +void launch_indexer(const FusedQIndexerParams& params, cudaStream_t stream) { + const uint32_t works = params.batch_size * params.num_heads; + const dim3 grid((works + kFusedQNumWarps - 1) / kFusedQNumWarps); + launch_kernel( + fused_q_indexer_rope_hadamard_quant_kernel, grid, kFusedQBlockSize, stream, UsePDL, params); +} + +void fused_q_norm_rope_cuda( + const at::Tensor& q_input, + const at::Tensor& q_output, + const at::Tensor& freqs_cis, + const at::Tensor& positions, + double eps) { + check_cuda_tensor(q_input, "q_input"); + check_same_device(q_input, q_output, "q_output"); + check_same_device(q_input, freqs_cis, "freqs_cis"); + check_same_device(q_input, positions, "positions"); + TORCH_CHECK(q_input.scalar_type() == at::kBFloat16, "q_input must be bfloat16"); + TORCH_CHECK(q_output.scalar_type() == at::kBFloat16, "q_output must be bfloat16"); + TORCH_CHECK(q_input.sizes() == q_output.sizes(), "q_input and q_output shapes must match"); + TORCH_CHECK(q_input.dim() == 3 && q_input.size(2) == kMainHeadDim, "q_input must be [B, H, 512]"); + TORCH_CHECK(q_input.stride(2) == 1 && q_input.stride(1) == kMainHeadDim, "q_input head rows must be contiguous"); + TORCH_CHECK( + q_output.stride(2) == 1 && q_output.stride(1) == kMainHeadDim, "q_output head rows must be contiguous"); + TORCH_CHECK( + freqs_cis.dim() == 2 && freqs_cis.size(1) == kMainRopeDim && freqs_cis.scalar_type() == at::kFloat && + freqs_cis.is_contiguous(), + "freqs_cis must be contiguous [max_pos, 64] float32"); + TORCH_CHECK( + positions.dim() == 1 && positions.size(0) == q_input.size(0) && positions.is_contiguous(), + "positions must be contiguous [B]"); + TORCH_CHECK( + positions.scalar_type() == at::kInt || positions.scalar_type() == at::kLong, + "positions must be int32 or int64"); + if (q_input.size(0) == 0) return; + + c10::cuda::CUDAGuard guard(q_input.device()); + const auto stream = at::cuda::getCurrentCUDAStream(); + const FusedQNormRopeParams params{ + reinterpret_cast(q_input.data_ptr()), + reinterpret_cast<__nv_bfloat16*>(q_output.data_ptr()), + freqs_cis.data_ptr(), + positions.data_ptr(), + q_input.stride(0), + q_output.stride(0), + static_cast(q_input.size(0)), + static_cast(q_input.size(1)), + static_cast(eps), + }; + const bool use_pdl = device_supports_pdl(); + if (positions.scalar_type() == at::kInt) { + use_pdl ? launch_q_norm(params, stream) : launch_q_norm(params, stream); + } else { + use_pdl ? launch_q_norm(params, stream) : launch_q_norm(params, stream); + } +} + +void fused_k_norm_rope_flashmla_cuda( + const at::Tensor& kv, + const at::Tensor& kv_weight, + const at::Tensor& freqs_cis, + const at::Tensor& positions, + const at::Tensor& out_loc, + const at::Tensor& kvcache, + double eps, + int64_t page_size) { + check_cuda_tensor(kv, "kv"); + check_same_device(kv, kv_weight, "kv_weight"); + check_same_device(kv, freqs_cis, "freqs_cis"); + check_same_device(kv, positions, "positions"); + check_same_device(kv, out_loc, "out_loc"); + check_same_device(kv, kvcache, "kvcache"); + TORCH_CHECK(kv.dim() == 2 && kv.size(1) == kMainHeadDim && kv.stride(1) == 1, "kv must be [B, 512]"); + TORCH_CHECK(kv.scalar_type() == at::kBFloat16, "kv must be bfloat16"); + TORCH_CHECK( + kv_weight.dim() == 1 && kv_weight.size(0) == kMainHeadDim && kv_weight.scalar_type() == at::kBFloat16 && + kv_weight.is_contiguous(), + "kv_weight must be contiguous [512] bfloat16"); + TORCH_CHECK( + freqs_cis.dim() == 2 && freqs_cis.size(1) == kMainRopeDim && freqs_cis.scalar_type() == at::kFloat && + freqs_cis.is_contiguous(), + "freqs_cis must be contiguous [max_pos, 64] float32"); + TORCH_CHECK( + positions.dim() == 1 && positions.size(0) == kv.size(0) && positions.is_contiguous(), + "positions must be contiguous [B]"); + TORCH_CHECK( + positions.scalar_type() == at::kInt || positions.scalar_type() == at::kLong, + "positions must be int32 or int64"); + TORCH_CHECK( + out_loc.dim() == 1 && out_loc.size(0) == kv.size(0) && out_loc.scalar_type() == at::kInt && + out_loc.is_contiguous(), + "out_loc must be contiguous [B] int32"); + TORCH_CHECK( + kvcache.dim() == 2 && kvcache.scalar_type() == at::kByte && kvcache.is_contiguous() && + kvcache.size(1) == kFlashMLAPageBytes, + "kvcache must be contiguous [num_pages, 74880] uint8"); + TORCH_CHECK(page_size == kFlashMLAPageSize, "DSV4 main K CUDA kernel requires page_size=128"); + if (kv.size(0) == 0) return; + + c10::cuda::CUDAGuard guard(kv.device()); + const auto stream = at::cuda::getCurrentCUDAStream(); + const FusedKNormRopeFlashMLAParams params{ + reinterpret_cast(kv.data_ptr()), + reinterpret_cast(kv_weight.data_ptr()), + freqs_cis.data_ptr(), + positions.data_ptr(), + out_loc.data_ptr(), + kvcache.data_ptr(), + kv.stride(0), + static_cast(kv.size(0)), + static_cast(eps), + }; + const bool use_pdl = device_supports_pdl(); + if (positions.scalar_type() == at::kInt) { + use_pdl ? launch_k_norm(params, stream) : launch_k_norm(params, stream); + } else { + use_pdl ? launch_k_norm(params, stream) : launch_k_norm(params, stream); + } +} + +void fused_q_indexer_rope_hadamard_quant_cuda( + const at::Tensor& q_input, + const at::Tensor& q_fp8, + const at::Tensor& weight, + const at::Tensor& weights_out, + double weight_scale, + const at::Tensor& freqs_cis, + const at::Tensor& positions) { + check_cuda_tensor(q_input, "q_input"); + check_same_device(q_input, q_fp8, "q_fp8"); + check_same_device(q_input, weight, "weight"); + check_same_device(q_input, weights_out, "weights_out"); + check_same_device(q_input, freqs_cis, "freqs_cis"); + check_same_device(q_input, positions, "positions"); + TORCH_CHECK( + q_input.dim() == 3 && q_input.size(2) == kIndexerHeadDim && q_input.scalar_type() == at::kBFloat16 && + q_input.is_contiguous(), + "q_input must be contiguous [B, H, 128] bfloat16"); + TORCH_CHECK( + q_fp8.sizes() == q_input.sizes() && q_fp8.scalar_type() == at::kFloat8_e4m3fn && q_fp8.is_contiguous(), + "q_fp8 must be contiguous [B, H, 128] float8_e4m3fn"); + TORCH_CHECK( + weight.dim() == 2 && weight.size(0) == q_input.size(0) && weight.size(1) == q_input.size(1) && + weight.scalar_type() == at::kBFloat16 && weight.is_contiguous(), + "weight must be contiguous [B, H] bfloat16"); + TORCH_CHECK( + weights_out.dim() == 3 && weights_out.size(0) == q_input.size(0) && + weights_out.size(1) == q_input.size(1) && weights_out.size(2) == 1 && + weights_out.scalar_type() == at::kFloat && weights_out.is_contiguous(), + "weights_out must be contiguous [B, H, 1] float32"); + TORCH_CHECK( + freqs_cis.dim() == 2 && freqs_cis.size(1) == kIndexerRopeDim && freqs_cis.scalar_type() == at::kFloat && + freqs_cis.is_contiguous(), + "freqs_cis must be contiguous [max_pos, 64] float32"); + TORCH_CHECK( + positions.dim() == 1 && positions.size(0) == q_input.size(0) && positions.is_contiguous(), + "positions must be contiguous [B]"); + TORCH_CHECK( + positions.scalar_type() == at::kInt || positions.scalar_type() == at::kLong, + "positions must be int32 or int64"); + if (q_input.size(0) == 0) return; + + c10::cuda::CUDAGuard guard(q_input.device()); + const auto stream = at::cuda::getCurrentCUDAStream(); + const FusedQIndexerParams params{ + reinterpret_cast(q_input.data_ptr()), + reinterpret_cast(q_fp8.data_ptr()), + reinterpret_cast(weight.data_ptr()), + weights_out.data_ptr(), + static_cast(weight_scale), + freqs_cis.data_ptr(), + positions.data_ptr(), + static_cast(q_input.size(0)), + static_cast(q_input.size(1)), + }; + const bool use_pdl = device_supports_pdl(); + if (positions.scalar_type() == at::kInt) { + use_pdl ? launch_indexer(params, stream) : launch_indexer(params, stream); + } else { + use_pdl ? launch_indexer(params, stream) : launch_indexer(params, stream); + } +} + +} // namespace + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { + module.def("fused_q_norm_rope", &fused_q_norm_rope_cuda, "DSV4 fused main-Q RMSNorm + RoPE"); + module.def( + "fused_k_norm_rope_flashmla", + &fused_k_norm_rope_flashmla_cuda, + "DSV4 fused main-K RMSNorm + RoPE + FlashMLA cache write"); + module.def( + "fused_q_indexer_rope_hadamard_quant", + &fused_q_indexer_rope_hadamard_quant_cuda, + "DSV4 fused indexer-Q RoPE + Hadamard + FP8 quant"); +} diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu new file mode 100644 index 0000000000..53d430243d --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu @@ -0,0 +1,345 @@ +// Copyright 2023-2024 SGLang Team +// SPDX-License-Identifier: Apache-2.0 +// +// DeepSeek-V4 c4-indexer top-k selection + page-translate. +// +// Adapted from SGLang commit 8cea0473ea5299bc04885f8f6ba71269415a39b5, +// python/sglang/jit_kernel/csrc/deepseek_v4/topk_v1.cuh. LightLLM replaces +// the tvm::ffi TensorView / TensorMatcher binding with a torch cpp_extension +// launcher while preserving the original Hopper PDL protocol. + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace { + +constexpr uint32_t kTopK = 512; +constexpr uint32_t kTopKBlockSize = 512; +constexpr uint32_t kSMEM = 16 * 1024 * sizeof(uint32_t); // 64KB (bytes) + +template +__device__ __forceinline__ void pdl_wait_primary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) asm volatile("griddepcontrol.wait;" ::: "memory"); +#endif +} + +template +__device__ __forceinline__ void pdl_trigger_secondary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) asm volatile("griddepcontrol.launch_dependents;" :::); +#endif +} + +__device__ __forceinline__ uint8_t convert_to_uint8(float x) { + __half h = __float2half_rn(x); + uint16_t bits = __half_as_ushort(h); + uint16_t key = (bits & 0x8000) ? static_cast(~bits) : static_cast(bits | 0x8000); + return static_cast(key >> 8); +} + +__device__ __forceinline__ uint32_t convert_to_uint32(float x) { + uint32_t bits = __float_as_uint(x); + return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); +} + +__device__ __forceinline__ int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, + uint32_t page_bits) { + const uint32_t mask = (1u << page_bits) - 1u; + return (page_table[i >> page_bits] << page_bits) | (i & mask); +} + +__device__ void naive_transform(const int32_t* __restrict__ page_table, int32_t* __restrict__ indices, + int32_t* __restrict__ raw_indices, const uint32_t length, + const uint32_t page_bits) { + if (const auto tx = threadIdx.x; tx < length) { + indices[tx] = page_to_indices(page_table, tx, page_bits); + if (raw_indices != nullptr) raw_indices[tx] = tx; + } else if (tx < kTopK) { + indices[tx] = -1; // fill invalid indices to -1 + if (raw_indices != nullptr) raw_indices[tx] = -1; + } +} + +__device__ void radix_topk(const float* __restrict__ input, int32_t* __restrict__ output, + const uint32_t length) { + constexpr uint32_t RADIX = 256; + constexpr uint32_t BLOCK_SIZE = kTopKBlockSize; + constexpr uint32_t SMEM_INPUT_SIZE = kSMEM / (2 * sizeof(int32_t)); + + alignas(128) __shared__ uint32_t _s_histogram_buf[2][RADIX + 32]; + alignas(128) __shared__ uint32_t s_counter; + alignas(128) __shared__ uint32_t s_threshold_bin_id; + alignas(128) __shared__ uint32_t s_num_input[2]; + alignas(128) __shared__ int32_t s_last_remain; + + extern __shared__ uint32_t s_input_idx[][kSMEM / (2 * sizeof(int32_t))]; + + const uint32_t tx = threadIdx.x; + uint32_t remain_topk = kTopK; + auto& s_histogram = _s_histogram_buf[0]; + + const auto run_cumsum = [&] { +#pragma unroll 8 + for (int32_t i = 0; i < 8; ++i) { + static_assert(1 << 8 == RADIX); + if (tx < RADIX) { + const auto j = 1 << i; + const auto k = i & 1; + auto value = _s_histogram_buf[k][tx]; + if (tx + j < RADIX) { + value += _s_histogram_buf[k][tx + j]; + } + _s_histogram_buf[k ^ 1][tx] = value; + } + __syncthreads(); + } + }; + + // stage 1: 8bit coarse histogram + if (tx < RADIX + 1) s_histogram[tx] = 0; + __syncthreads(); + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const auto bin = convert_to_uint8(input[idx]); + ::atomicAdd(&s_histogram[bin], 1); + } + __syncthreads(); + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { + s_threshold_bin_id = tx; + s_num_input[0] = 0; + s_counter = 0; + } + __syncthreads(); + + const auto threshold_bin = s_threshold_bin_id; + remain_topk -= s_histogram[threshold_bin + 1]; + if (remain_topk == 0) { + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const uint32_t bin = convert_to_uint8(input[idx]); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } + } + __syncthreads(); + return; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + + for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { + const float raw_input = input[idx]; + const uint32_t bin = convert_to_uint8(raw_input); + if (bin > threshold_bin) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } else if (bin == threshold_bin) { + const auto pos = ::atomicAdd(&s_num_input[0], 1); + if (pos < SMEM_INPUT_SIZE) { + s_input_idx[0][pos] = idx; + const auto sbin = convert_to_uint32(raw_input); + const auto sub_bin = (sbin >> 24) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + __syncthreads(); + } + + // stage 2: refine with 8bit radix passes +#pragma unroll 4 + for (int round = 0; round < 4; ++round) { + const auto r_idx = round % 2; + + const auto raw_num_input = s_num_input[r_idx]; + const auto num_input = raw_num_input < SMEM_INPUT_SIZE ? raw_num_input : SMEM_INPUT_SIZE; + + run_cumsum(); + if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { + s_threshold_bin_id = tx; + s_num_input[r_idx ^ 1] = 0; + s_last_remain = remain_topk - s_histogram[tx + 1]; + } + __syncthreads(); + + const auto threshold_bin2 = s_threshold_bin_id; + remain_topk -= s_histogram[threshold_bin2 + 1]; + + if (remain_topk == 0) { + for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(input[idx]) >> offset) & 0xFF; + if (bin > threshold_bin2) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } + } + __syncthreads(); + break; + } else { + __syncthreads(); + if (tx < RADIX + 1) { + s_histogram[tx] = 0; + } + __syncthreads(); + for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { + const auto idx = s_input_idx[r_idx][i]; + const auto raw_input = input[idx]; + const auto offset = 24 - round * 8; + const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF; + if (bin > threshold_bin2) { + const auto pos = ::atomicAdd(&s_counter, 1); + output[pos] = idx; + } else if (bin == threshold_bin2) { + if (round == 3) { + const auto pos = ::atomicAdd(&s_last_remain, -1); + if (pos > 0) { + output[kTopK - pos] = idx; + } + } else { + const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1); + if (pos < SMEM_INPUT_SIZE) { + s_input_idx[r_idx ^ 1][pos] = idx; + const auto sbin = convert_to_uint32(raw_input); + const auto sub_bin = (sbin >> (offset - 8)) & 0xFF; + ::atomicAdd(&s_histogram[sub_bin], 1); + } + } + } + } + __syncthreads(); + } + } +} + +struct TopKParams { + const float* scores; + const int32_t* seq_lens; + const int32_t* page_table; + int32_t* page_indices; + int32_t* raw_indices; + int64_t score_stride; + int64_t page_table_stride; + uint32_t page_bits; +}; + +template +__global__ void topk_transform_kernel(const __grid_constant__ TopKParams params) { + const uint32_t work_id = blockIdx.x; + const uint32_t seq_len = params.seq_lens[work_id]; + const auto score_ptr = params.scores + work_id * params.score_stride; + const auto page_ptr = params.page_table + work_id * params.page_table_stride; + const auto indices_ptr = params.page_indices + work_id * kTopK; + const auto raw_indices_ptr = params.raw_indices != nullptr ? params.raw_indices + work_id * kTopK : nullptr; + const uint32_t page_bits = params.page_bits; + + pdl_wait_primary(); + + if (seq_len <= kTopK) { + naive_transform(page_ptr, indices_ptr, raw_indices_ptr, seq_len, page_bits); + } else { + __shared__ int32_t s_topk_indices[kTopK]; + radix_topk(score_ptr, s_topk_indices, seq_len); + const auto tx = threadIdx.x; + indices_ptr[tx] = page_to_indices(page_ptr, s_topk_indices[tx], page_bits); + if (raw_indices_ptr != nullptr) raw_indices_ptr[tx] = s_topk_indices[tx]; + } + + pdl_trigger_secondary(); +} + +template +void launch_topk(const TopKParams& params, uint32_t batch_size, cudaStream_t stream) { + constexpr uint32_t smem = kSMEM + sizeof(int32_t); + static const cudaError_t smem_result = cudaFuncSetAttribute( + topk_transform_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); + C10_CUDA_CHECK(smem_result); + + cudaLaunchConfig_t config{}; + config.gridDim = batch_size; + config.blockDim = kTopKBlockSize; + config.dynamicSmemBytes = smem; + config.stream = stream; + + cudaLaunchAttribute attribute{}; + if constexpr (UsePDL) { + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = true; + config.attrs = &attribute; + config.numAttrs = 1; + } + C10_CUDA_CHECK(cudaLaunchKernelEx(&config, topk_transform_kernel, params)); +} + +} // namespace + +void topk_transform_512_cuda(at::Tensor scores, at::Tensor seq_lens, at::Tensor page_table, + at::Tensor page_indices, int64_t page_size, c10::optional raw_indices) { + TORCH_CHECK(scores.is_cuda(), "scores must be a CUDA tensor"); + TORCH_CHECK( + seq_lens.is_cuda() && page_table.is_cuda() && page_indices.is_cuda() && + seq_lens.get_device() == scores.get_device() && page_table.get_device() == scores.get_device() && + page_indices.get_device() == scores.get_device(), + "all tensors must be on the same CUDA device"); + TORCH_CHECK(scores.dim() == 2 && scores.dtype() == at::kFloat, "scores must be [B, S] float32"); + TORCH_CHECK( + seq_lens.dim() == 1 && seq_lens.size(0) == scores.size(0) && seq_lens.dtype() == at::kInt && + seq_lens.is_contiguous(), + "seq_lens must be [B] int32 contiguous"); + TORCH_CHECK( + page_table.dim() == 2 && page_table.size(0) == scores.size(0) && page_table.dtype() == at::kInt, + "page_table must be [B, P] int32"); + TORCH_CHECK(page_indices.dim() == 2 && page_indices.dtype() == at::kInt && page_indices.is_contiguous(), + "page_indices must be [B, 512] int32 contiguous"); + TORCH_CHECK(page_indices.size(0) == scores.size(0), "page_indices first dim must match scores"); + TORCH_CHECK(page_indices.size(1) == (int64_t)kTopK, "page_indices second dim must be 512"); + TORCH_CHECK(scores.stride(1) == 1 && page_table.stride(1) == 1, "scores/page_table last dim must be contiguous"); + TORCH_CHECK(page_size > 0 && (page_size & (page_size - 1)) == 0, "page_size must be a power of two"); + + const uint32_t batch_size = scores.size(0); + if (batch_size == 0) return; + const uint32_t page_bits = __builtin_ctzll(page_size); + + int32_t* raw_ptr = nullptr; + if (raw_indices.has_value()) { + auto& r = raw_indices.value(); + TORCH_CHECK( + r.is_cuda() && r.get_device() == scores.get_device() && r.dim() == 2 && r.size(0) == scores.size(0) && + r.size(1) == (int64_t)kTopK && r.dtype() == at::kInt && r.is_contiguous(), + "raw_indices must be [B, 512] int32 contiguous on the same CUDA device"); + raw_ptr = r.data_ptr(); + } + + c10::cuda::CUDAGuard guard(scores.device()); + auto stream = at::cuda::getCurrentCUDAStream(); + const TopKParams params{ + scores.data_ptr(), + seq_lens.data_ptr(), + page_table.data_ptr(), + page_indices.data_ptr(), + raw_ptr, + scores.stride(0), + page_table.stride(0), + page_bits, + }; + if (at::cuda::getCurrentDeviceProperties()->major >= 9) { + launch_topk(params, batch_size, stream); + } else { + launch_topk(params, batch_size, stream); + } +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("topk_transform_512", &topk_transform_512_cuda, "DeepSeek-V4 c4 indexer top-512 + page translate"); +} diff --git a/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py b/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py new file mode 100644 index 0000000000..0ef4434f2e --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py @@ -0,0 +1,98 @@ +import functools +import hashlib +import os + +import torch + + +@functools.lru_cache(maxsize=1) +def _load_cuda(): + from torch.utils.cpp_extension import load + + src = os.path.join(os.path.dirname(__file__), "csrc", "norm_rope.cu") + flags = ["-O3"] + with open(src, "rb") as source_file: + source = source_file.read() + capability = torch.cuda.get_device_capability() + cache_key = b"\0".join( + [ + source, + " ".join(flags).encode(), + torch.__version__.encode(), + str(torch.version.cuda).encode(), + f"sm{capability[0]}{capability[1]}".encode(), + os.environ.get("TORCH_CUDA_ARCH_LIST", "").encode(), + ] + ) + module_name = f"lightllm_dsv4_norm_rope_v1_{hashlib.sha256(cache_key).hexdigest()[:16]}" + return load( + name=module_name, + sources=[src], + extra_cuda_cflags=flags, + verbose=False, + ) + + +def _as_interleaved_freqs(freqs_cis: torch.Tensor) -> torch.Tensor: + assert freqs_cis.dtype == torch.complex64 + freqs_real = torch.view_as_real(freqs_cis).flatten(-2) + assert freqs_real.is_contiguous() + return freqs_real + + +@torch.no_grad() +def fused_q_norm_rope( + q_input: torch.Tensor, + q_output: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +) -> None: + _load_cuda().fused_q_norm_rope(q_input, q_output, _as_interleaved_freqs(freqs_cis), positions, float(eps)) + return + + +@torch.no_grad() +def fused_k_norm_rope_flashmla( + kv: torch.Tensor, + kv_weight: torch.Tensor, + eps: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, + out_loc: torch.Tensor, + kvcache: torch.Tensor, + page_size: int, +) -> None: + _load_cuda().fused_k_norm_rope_flashmla( + kv, + kv_weight, + _as_interleaved_freqs(freqs_cis), + positions, + out_loc, + kvcache, + float(eps), + int(page_size), + ) + return + + +@torch.no_grad() +def fused_q_indexer_rope_hadamard_quant( + q_input: torch.Tensor, + weight: torch.Tensor, + weight_scale: float, + freqs_cis: torch.Tensor, + positions: torch.Tensor, +): + q_fp8 = torch.empty(q_input.shape, dtype=torch.float8_e4m3fn, device=q_input.device) + weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) + _load_cuda().fused_q_indexer_rope_hadamard_quant( + q_input, + q_fp8, + weight, + weights_out, + float(weight_scale), + _as_interleaved_freqs(freqs_cis), + positions, + ) + return q_fp8, weights_out diff --git a/lightllm/models/deepseek_v4/triton_kernel/topk_transform.py b/lightllm/models/deepseek_v4/triton_kernel/topk_transform.py new file mode 100644 index 0000000000..76256f5765 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/topk_transform.py @@ -0,0 +1,61 @@ +import functools +import hashlib +import os + +import torch + + +@functools.lru_cache(maxsize=1) +def _load_cuda(): + from torch.utils.cpp_extension import load + + src = os.path.join(os.path.dirname(__file__), "csrc", "topk_transform.cu") + flags = ["-O3"] + with open(src, "rb") as source_file: + source = source_file.read() + capability = torch.cuda.get_device_capability() + cache_key = b"\0".join( + [ + source, + " ".join(flags).encode(), + torch.__version__.encode(), + str(torch.version.cuda).encode(), + f"sm{capability[0]}{capability[1]}".encode(), + os.environ.get("TORCH_CUDA_ARCH_LIST", "").encode(), + ] + ) + module_name = f"lightllm_dsv4_topk_v1_{hashlib.sha256(cache_key).hexdigest()[:16]}" + return load( + name=module_name, + sources=[src], + extra_cuda_cflags=flags, + verbose=False, + ) + + +@torch.no_grad() +def topk_transform_512( + scores: torch.Tensor, + seq_lens: torch.Tensor, + page_tables: torch.Tensor, + out_page_indices: torch.Tensor, + page_size: int, + out_raw_indices: torch.Tensor = None, +) -> None: + """Masked top-512 selection over per-token scores + page-translate, for the DeepSeek-V4 c4 + indexer. Drop-in replacement for the former vendored topk_transform_512 op. + + LightLLM-local CUDA radix port (from SGLang topk_v1.cuh, TVM-FFI stripped and PDL preserved): it + early-exits per token at seq_len, so it matches the original perf (unlike torch.topk which + scans the full captured c4_cap width). Output is an unordered SET of physical c4 slots (-1 pad). + + Args: + scores: [T, c4_cap] fp32 (deep_gemm fp8_paged_mqa_logits output, -inf beyond ctx) + seq_lens: [T] int32 (valid causal entries per token) + page_tables: [T, npages] int32 (logical->physical c4 page map) + out_page_indices: [T, 512] int32 (output physical slots, -1 pad) + page_size: c4 pool page size (64) + out_raw_indices: optional [T, 512] int32 (raw logical indices, -1 pad) + """ + _load_cuda().topk_transform_512(scores, seq_lens, page_tables, out_page_indices, int(page_size), out_raw_indices) + return diff --git a/lightllm/third_party/__init__.py b/lightllm/third_party/__init__.py deleted file mode 100644 index 2adb50db25..0000000000 --- a/lightllm/third_party/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Third-party source subsets vendored for LightLLM runtime support.""" diff --git a/lightllm/third_party/sglang_jit/LICENSE b/lightllm/third_party/sglang_jit/LICENSE deleted file mode 100755 index 9c422689c8..0000000000 --- a/lightllm/third_party/sglang_jit/LICENSE +++ /dev/null @@ -1,201 +0,0 @@ - Apache License - Version 2.0, January 2004 - http://www.apache.org/licenses/ - - TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION - - 1. Definitions. - - "License" shall mean the terms and conditions for use, reproduction, - and distribution as defined by Sections 1 through 9 of this document. - - "Licensor" shall mean the copyright owner or entity authorized by - the copyright owner that is granting the License. - - "Legal Entity" shall mean the union of the acting entity and all - other entities that control, are controlled by, or are under common - control with that entity. For the purposes of this definition, - "control" means (i) the power, direct or indirect, to cause the - direction or management of such entity, whether by contract or - otherwise, or (ii) ownership of fifty percent (50%) or more of the - outstanding shares, or (iii) beneficial ownership of such entity. - - "You" (or "Your") shall mean an individual or Legal Entity - exercising permissions granted by this License. - - "Source" form shall mean the preferred form for making modifications, - including but not limited to software source code, documentation - source, and configuration files. - - "Object" form shall mean any form resulting from mechanical - transformation or translation of a Source form, including but - not limited to compiled object code, generated documentation, - and conversions to other media types. - - "Work" shall mean the work of authorship, whether in Source or - Object form, made available under the License, as indicated by a - copyright notice that is included in or attached to the work - (an example is provided in the Appendix below). - - "Derivative Works" shall mean any work, whether in Source or Object - form, that is based on (or derived from) the Work and for which the - editorial revisions, annotations, elaborations, or other modifications - represent, as a whole, an original work of authorship. For the purposes - of this License, Derivative Works shall not include works that remain - separable from, or merely link (or bind by name) to the interfaces of, - the Work and Derivative Works thereof. - - "Contribution" shall mean any work of authorship, including - the original version of the Work and any modifications or additions - to that Work or Derivative Works thereof, that is intentionally - submitted to Licensor for inclusion in the Work by the copyright owner - or by an individual or Legal Entity authorized to submit on behalf of - the copyright owner. For the purposes of this definition, "submitted" - means any form of electronic, verbal, or written communication sent - to the Licensor or its representatives, including but not limited to - communication on electronic mailing lists, source code control systems, - and issue tracking systems that are managed by, or on behalf of, the - Licensor for the purpose of discussing and improving the Work, but - excluding communication that is conspicuously marked or otherwise - designated in writing by the copyright owner as "Not a Contribution." - - "Contributor" shall mean Licensor and any individual or Legal Entity - on behalf of whom a Contribution has been received by Licensor and - subsequently incorporated within the Work. - - 2. Grant of Copyright License. Subject to the terms and conditions of - this License, each Contributor hereby grants to You a perpetual, - worldwide, non-exclusive, no-charge, royalty-free, irrevocable - copyright license to reproduce, prepare Derivative Works of, - publicly display, publicly perform, sublicense, and distribute the - Work and such Derivative Works in Source or Object form. - - 3. Grant of Patent License. Subject to the terms and conditions of - this License, each Contributor hereby grants to You a perpetual, - worldwide, non-exclusive, no-charge, royalty-free, irrevocable - (except as stated in this section) patent license to make, have made, - use, offer to sell, sell, import, and otherwise transfer the Work, - where such license applies only to those patent claims licensable - by such Contributor that are necessarily infringed by their - Contribution(s) alone or by combination of their Contribution(s) - with the Work to which such Contribution(s) was submitted. If You - institute patent litigation against any entity (including a - cross-claim or counterclaim in a lawsuit) alleging that the Work - or a Contribution incorporated within the Work constitutes direct - or contributory patent infringement, then any patent licenses - granted to You under this License for that Work shall terminate - as of the date such litigation is filed. - - 4. Redistribution. You may reproduce and distribute copies of the - Work or Derivative Works thereof in any medium, with or without - modifications, and in Source or Object form, provided that You - meet the following conditions: - - (a) You must give any other recipients of the Work or - Derivative Works a copy of this License; and - - (b) You must cause any modified files to carry prominent notices - stating that You changed the files; and - - (c) You must retain, in the Source form of any Derivative Works - that You distribute, all copyright, patent, trademark, and - attribution notices from the Source form of the Work, - excluding those notices that do not pertain to any part of - the Derivative Works; and - - (d) If the Work includes a "NOTICE" text file as part of its - distribution, then any Derivative Works that You distribute must - include a readable copy of the attribution notices contained - within such NOTICE file, excluding those notices that do not - pertain to any part of the Derivative Works, in at least one - of the following places: within a NOTICE text file distributed - as part of the Derivative Works; within the Source form or - documentation, if provided along with the Derivative Works; or, - within a display generated by the Derivative Works, if and - wherever such third-party notices normally appear. The contents - of the NOTICE file are for informational purposes only and - do not modify the License. You may add Your own attribution - notices within Derivative Works that You distribute, alongside - or as an addendum to the NOTICE text from the Work, provided - that such additional attribution notices cannot be construed - as modifying the License. - - You may add Your own copyright statement to Your modifications and - may provide additional or different license terms and conditions - for use, reproduction, or distribution of Your modifications, or - for any such Derivative Works as a whole, provided Your use, - reproduction, and distribution of the Work otherwise complies with - the conditions stated in this License. - - 5. Submission of Contributions. Unless You explicitly state otherwise, - any Contribution intentionally submitted for inclusion in the Work - by You to the Licensor shall be under the terms and conditions of - this License, without any additional terms or conditions. - Notwithstanding the above, nothing herein shall supersede or modify - the terms of any separate license agreement you may have executed - with Licensor regarding such Contributions. - - 6. Trademarks. This License does not grant permission to use the trade - names, trademarks, service marks, or product names of the Licensor, - except as required for reasonable and customary use in describing the - origin of the Work and reproducing the content of the NOTICE file. - - 7. Disclaimer of Warranty. Unless required by applicable law or - agreed to in writing, Licensor provides the Work (and each - Contributor provides its Contributions) on an "AS IS" BASIS, - WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or - implied, including, without limitation, any warranties or conditions - of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A - PARTICULAR PURPOSE. You are solely responsible for determining the - appropriateness of using or redistributing the Work and assume any - risks associated with Your exercise of permissions under this License. - - 8. Limitation of Liability. In no event and under no legal theory, - whether in tort (including negligence), contract, or otherwise, - unless required by applicable law (such as deliberate and grossly - negligent acts) or agreed to in writing, shall any Contributor be - liable to You for damages, including any direct, indirect, special, - incidental, or consequential damages of any character arising as a - result of this License or out of the use or inability to use the - Work (including but not limited to damages for loss of goodwill, - work stoppage, computer failure or malfunction, or any and all - other commercial damages or losses), even if such Contributor - has been advised of the possibility of such damages. - - 9. Accepting Warranty or Additional Liability. While redistributing - the Work or Derivative Works thereof, You may choose to offer, - and charge a fee for, acceptance of support, warranty, indemnity, - or other liability obligations and/or rights consistent with this - License. However, in accepting such obligations, You may act only - on Your own behalf and on Your sole responsibility, not on behalf - of any other Contributor, and only if You agree to indemnify, - defend, and hold each Contributor harmless for any liability - incurred by, or claims asserted against, such Contributor by reason - of your accepting any such warranty or additional liability. - - END OF TERMS AND CONDITIONS - - APPENDIX: How to apply the Apache License to your work. - - To apply the Apache License to your work, attach the following - boilerplate notice, with the fields enclosed by brackets "[]" - replaced with your own identifying information. (Don't include - the brackets!) The text should be enclosed in the appropriate - comment syntax for the file format. We also recommend that a - file or class name and description of purpose be included on the - same "printed page" as the copyright notice for easier - identification within third-party archives. - - Copyright 2023-2024 SGLang Team - - Licensed under the Apache License, Version 2.0 (the "License"); - you may not use this file except in compliance with the License. - You may obtain a copy of the License at - - http://www.apache.org/licenses/LICENSE-2.0 - - Unless required by applicable law or agreed to in writing, software - distributed under the License is distributed on an "AS IS" BASIS, - WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - See the License for the specific language governing permissions and - limitations under the License. diff --git a/lightllm/third_party/sglang_jit/README.md b/lightllm/third_party/sglang_jit/README.md deleted file mode 100644 index 4f68c9cfd8..0000000000 --- a/lightllm/third_party/sglang_jit/README.md +++ /dev/null @@ -1,13 +0,0 @@ -# Vendored SGLang JIT Subset - -This directory contains the minimal SGLang JIT source subset needed by the -DeepSeek-V4 LightLLM implementation. - -Source: https://github.com/sgl-project/sglang -Commit: 8cea0473ea5299bc04885f8f6ba71269415a39b5 -License: Apache License 2.0, copied in `LICENSE`. - -Local changes: -- The Python imports were moved from `sglang.jit_kernel.*` to - `lightllm.third_party.sglang_jit.*`. -- The package exports only the DSv4 functions used by LightLLM. diff --git a/lightllm/third_party/sglang_jit/__init__.py b/lightllm/third_party/sglang_jit/__init__.py deleted file mode 100644 index 164d545b4e..0000000000 --- a/lightllm/third_party/sglang_jit/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Vendored SGLang JIT kernels used by DeepSeek-V4.""" diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh deleted file mode 100644 index 3a89e8114c..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128.cuh +++ /dev/null @@ -1,522 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include - -namespace { - -using Plan128 = device::compress::PrefillPlan; -using IndiceT = int32_t; - -/// \brief Each thread will handle this many elements (split along head_dim) -constexpr int32_t kTileElements = 2; -/// \brief Each warp will handle this many elements (split along 128) -constexpr int32_t kElementsPerWarp = 8; -constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; -constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps; - -/// \brief Need to reduce register usage to increase occupancy -#define C128_KERNEL __global__ __launch_bounds__(kBlockSize, 2) - -struct Compress128DecodeParams { - /** - * \brief Shape: `[num_indices, 128, head_dim * 2]` \n - * last dimension layout: - * | kv current | score current | - */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[batch_size, head_dim * 2]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[batch_size, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[128, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, ]` */ - const IndiceT* __restrict__ seq_lens; - /** \NOTE: `batch_size` <= `num_indices` */ - uint32_t batch_size; -}; - -struct Compress128PrefillParams { - /** - * \brief Shape: `[num_indices, 128, head_dim * 2]` \n - * last dimension layout: - * | kv current | score current | - */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[batch_size, head_dim * 2]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[batch_size, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[128, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, ]`*/ - const int32_t* __restrict__ load_indices; - /** \brief The following part is plan info. */ - const Plan128* __restrict__ compress_plan; - const Plan128* __restrict__ write_plan; - uint32_t num_compress; - uint32_t num_write; -}; - -struct Compress128SharedBuffer { - using Storage = device::AlignedVector; - Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict - SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { - return data[warp_id][lane_id]; - } - SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { - return data[warp_id][lane_id][tile_id]; - } -}; - -template -SGL_DEVICE void c128_write( - T* kv_score_buf, // - const T* kv_score_src, - const int64_t head_dim, - const int32_t write_pos, - const uint32_t lane_id) { - using namespace device; - - using Storage = AlignedVector; - const auto element_size = head_dim * 2; - const auto gmem = tile::Memory{lane_id, kWarpThreads}; - kv_score_buf += write_pos * element_size; - - /// NOTE: Layout | [0] = kv | [1] = score | - Storage kv_score[2]; -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - kv_score[i] = gmem.load(kv_score_src + head_dim * i); - } -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - gmem.store(kv_score_buf + head_dim * i, kv_score[i]); - } -} - -template -SGL_DEVICE void c128_forward( - const InFloat* kv_score_buf, - const InFloat* kv_score_src, - OutFloat* kv_out, - const InFloat* score_bias, - const int64_t head_dim, - const int32_t window_len, - const uint32_t warp_id, - const uint32_t lane_id) { - using namespace device; - - const auto element_size = head_dim * 2; - const auto score_offset = head_dim; - - /// NOTE: part 1: load kv + score - using StorageIn = AlignedVector; - const auto gmem_in = tile::Memory{lane_id, kWarpThreads}; - StorageIn kv[kElementsPerWarp]; - StorageIn score[kElementsPerWarp]; - StorageIn bias[kElementsPerWarp]; - const int32_t warp_offset = warp_id * kElementsPerWarp; - -#pragma unroll - for (int32_t i = 0; i < 8; ++i) { - const int32_t j = i + warp_offset; - bias[i] = gmem_in.load(score_bias + j * head_dim); - } - -#pragma unroll - for (int32_t i = 0; i < kElementsPerWarp; ++i) { - const int32_t j = i + warp_offset; - const InFloat* src; - __builtin_assume(j < 128); - if (j < window_len) { - src = kv_score_buf + j * element_size; - } else { - /// NOTE: k in [-127, 0]. We'll load from the ragged `kv_score_src` - const int32_t k = j - 127; - src = kv_score_src + k * element_size; - } - kv[i] = gmem_in.load(src); - score[i] = gmem_in.load(src + score_offset); - } - - /// NOTE: part 2: safe online softmax + weighted sum - using TmpStorage = typename Compress128SharedBuffer::Storage; - __shared__ Compress128SharedBuffer s_local_val_max; - __shared__ Compress128SharedBuffer s_local_exp_sum; - __shared__ Compress128SharedBuffer s_local_product; - - TmpStorage tmp_val_max; - TmpStorage tmp_exp_sum; - TmpStorage tmp_product; - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - float score_fp32[kElementsPerWarp]; - -#pragma unroll - for (int32_t j = 0; j < kElementsPerWarp; ++j) { - score_fp32[j] = cast(score[j][i]) + cast(bias[j][i]); - } - - float max_value = score_fp32[0]; - float sum_exp_value = 0.0f; - -#pragma unroll - for (int32_t j = 1; j < kElementsPerWarp; ++j) { - const auto fp32_score = score_fp32[j]; - max_value = fmaxf(max_value, fp32_score); - } - - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - const auto fp32_score = score_fp32[j]; - const auto exp_score = expf(fp32_score - max_value); - sum_product += cast(kv[j][i]) * exp_score; - sum_exp_value += exp_score; - } - - tmp_val_max[i] = max_value; - tmp_exp_sum[i] = sum_exp_value; - tmp_product[i] = sum_product; - } - - // naturally aligned, so no bank conflict - s_local_val_max(warp_id, lane_id) = tmp_val_max; - s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; - s_local_product(warp_id, lane_id) = tmp_product; - - __syncthreads(); - - /// NOTE: part 3: online softmax - /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce - /// each reduce will consume `kNumWarps` threads (use partial warp reduction) - constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; - constexpr uint32_t kIteration = kReductionCount / kBlockSize; - -#pragma unroll - for (uint32_t i = 0; i < kIteration; ++i) { - /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` - const uint32_t j = i * kBlockSize + warp_id * kWarpThreads + lane_id; - /// NOTE: Range `[0, kNumWarps)` - const uint32_t local_warp_id = j % kNumWarps; - /// NOTE: Range `[0, kTileElements * kWarpThreads)` - const uint32_t local_elem_id = j / kNumWarps; - /// NOTE: Range `[0, kTileElements)` - const uint32_t local_tile_id = local_elem_id % kTileElements; - /// NOTE: Range `[0, kWarpThreads)` - const uint32_t local_lane_id = local_elem_id / kTileElements; - /// NOTE: each warp will access the whole tile (all `kTileElements`) - /// and for different lanes, the memory access only differ in `local_warp_id` - /// so there's no bank conflict in shared memory access. - static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); - const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); - const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); - const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); - const auto global_val_max = warp::reduce_max(local_val_max); - const auto rescale = expf(local_val_max - global_val_max); - const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); - const auto final_scale = rescale / global_exp_sum; - const auto global_product = warp::reduce_sum(local_product * final_scale); - kv_out[local_elem_id] = cast(global_product); - } -} - -template -C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - constexpr int64_t kElementSize = kHeadDim * 2; - static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); - - const auto& [ - _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score - indices, seq_lens, batch_size // decode info - ] = params; - const uint32_t warp_id = threadIdx.x / kWarpThreads; - const uint32_t lane_id = threadIdx.x % kWarpThreads; - - const uint32_t global_bid = blockIdx.x / kNumSplit; // batch id - const uint32_t global_sid = blockIdx.x % kNumSplit; // split id - if (global_bid >= batch_size) return; - - const int32_t index = indices[global_bid]; - const int32_t seq_len = seq_lens[global_bid]; - const int64_t split_offset = global_sid * kTileDim; - - // kv score - const auto kv_score_buffer = static_cast(_kv_score_buffer); - const auto kv_buf = kv_score_buffer + index * (kElementSize * 128) + split_offset; - - // kv input - const auto kv_score_input = static_cast(_kv_score_input); - const auto kv_src = kv_score_input + global_bid * kElementSize + split_offset; - - // kv output - const auto kv_compressed_output = static_cast(_kv_compressed_output); - const auto kv_out = kv_compressed_output + global_bid * kHeadDim + split_offset; - - // score bias (ape) - const auto score_bias = static_cast(_score_bias) + split_offset; - - PDLWaitPrimary(); - - /// NOTE: the write must be visible to the subsequent c128_forward, - /// so only the last warp can write to HBM - /// In addition, `position` = `seq_len - 1`. To avoid underflow, we use `seq_len + 127` - if (warp_id == kNumWarps - 1) { - c128_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 127) % 128, lane_id); - } - if (seq_len % 128 == 0) { - c128_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, /*window_len=*/128, warp_id, lane_id); - } - - PDLTriggerSecondary(); -} - -// compress kernel -template -C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - constexpr int64_t kElementSize = kHeadDim * 2; - static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); - - const auto& [ - _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score - indices, load_indices, compress_plan, write_plan, num_compress, num_write // prefill plan - ] = params; - const uint32_t warp_id = threadIdx.x / kWarpThreads; - const uint32_t lane_id = threadIdx.x % kWarpThreads; - - uint32_t global_id; - if constexpr (kWrite) { - // for write kernel, we use global warp_id to dispatch work - global_id = (blockIdx.x * blockDim.x + threadIdx.x) / kWarpThreads; - } else { - // for compress kernel, we use block id to dispatch work - global_id = blockIdx.x; // block id - } - const uint32_t global_pid = global_id / kNumSplit; // plan id - const uint32_t global_sid = global_id % kNumSplit; // split id - - /// NOTE: compiler can optimize this if-else at compile time - const auto num_plans = kWrite ? num_write : num_compress; - const auto plan_ptr = kWrite ? write_plan : compress_plan; - if (global_pid >= num_plans) return; - - const auto& [ragged_id, global_bid, position, window_len] = plan_ptr[global_pid]; - const auto indices_ptr = kWrite ? indices : load_indices; - - const int64_t split_offset = global_sid * kTileDim; - - // kv input - const auto kv_score_input = static_cast(_kv_score_input); - const auto kv_src = kv_score_input + ragged_id * kElementSize + split_offset; - - // kv output - const auto kv_compressed_output = static_cast(_kv_compressed_output); - const auto kv_out = kv_compressed_output + ragged_id * kHeadDim + split_offset; - - // score bias (ape) - const auto score_bias = static_cast(_score_bias) + split_offset; - - if (ragged_id == 0xFFFFFFFF) [[unlikely]] - return; - - const int32_t index = indices_ptr[global_bid]; - // kv score - const auto kv_score_buffer = static_cast(_kv_score_buffer); - const auto kv_buf = kv_score_buffer + index * (kElementSize * 128) + split_offset; - - PDLWaitPrimary(); - - // only responsible for the compress part - if constexpr (kWrite) { - c128_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 128, lane_id); - } else { - c128_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, window_len, warp_id, lane_id); - } - - PDLTriggerSecondary(); -} - -template -struct FlashCompress128Kernel { - static constexpr auto decode_kernel = flash_c128_decode; - template - static constexpr auto prefill_kernel = flash_c128_prefill; - static constexpr auto prefill_c_kernel = prefill_kernel; - static constexpr auto prefill_w_kernel = prefill_kernel; - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static constexpr uint32_t kWriteBlockSize = 128; - static constexpr uint32_t kWarpsPerWriteBlock = kWriteBlockSize / device::kWarpThreads; - - static void run_decode( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::Optional /* UNUSED */) { - using namespace host; - - // this should not happen in practice - auto B = SymbolicSize{"batch_size"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({-1, 128, kHeadDim * 2}) // kv score - .with_dtype() - .with_device(device) - .verify(kv_score_buffer); - TensorMatcher({B, kHeadDim * 2}) // kv score input - .with_dtype() - .with_device(device) - .verify(kv_score_input); - TensorMatcher({B, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device) - .verify(ape); - TensorMatcher({B}) // indices - .with_dtype() - .with_device(device) - .verify(indices); - TensorMatcher({B}) // seq lens - .with_dtype() - .with_device(device) - .verify(seq_lens); - - const auto batch_size = static_cast(B.unwrap()); - const auto params = Compress128DecodeParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .batch_size = batch_size, - }; - - const uint32_t num_blocks = batch_size * kNumSplit; - LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // - .enable_pdl(kUsePDL)(decode_kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView compress_plan, - const tvm::ffi::TensorView write_plan, - const tvm::ffi::Optional extra) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto N = SymbolicSize{"num_q_tokens"}; - auto X = SymbolicSize{"compress_tokens"}; - auto Y = SymbolicSize{"write_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 128, kHeadDim * 2}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({N, kHeadDim * 2}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({N, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - TensorMatcher({B}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - TensorMatcher({X, compress::kPrefillPlanDim}) // compress plan - .with_dtype() - .with_device(device_) - .verify(compress_plan); - TensorMatcher({Y, compress::kPrefillPlanDim}) // write plan - .with_dtype() - .with_device(device_) - .verify(write_plan); - - // might be needed for prefill write - const auto load_indices = extra.value_or(indices); - TensorMatcher({B}) // [read_positions] - .with_dtype() - .with_device(device_) - .verify(load_indices); - - const auto device = device_.unwrap(); - const auto batch_size = static_cast(B.unwrap()); - const auto num_q_tokens = static_cast(N.unwrap()); - const auto num_c = static_cast(X.unwrap()); - const auto num_w = static_cast(Y.unwrap()); - const auto params = Compress128PrefillParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .load_indices = static_cast(load_indices.data_ptr()), - .compress_plan = static_cast(compress_plan.data_ptr()), - .write_plan = static_cast(write_plan.data_ptr()), - .num_compress = num_c, - .num_write = num_w, - }; - RuntimeCheck(num_q_tokens >= batch_size, "num_q_tokens must be >= batch_size"); - RuntimeCheck(num_q_tokens >= std::max(num_c, num_w), "invalid prefill plan"); - - constexpr auto kBlockSize_C = kBlockSize; - constexpr auto kBlockSize_W = kWriteBlockSize; - if (const auto num_c_blocks = num_c * kNumSplit) { - LaunchKernel(num_c_blocks, kBlockSize_C, device) // - .enable_pdl(kUsePDL)(prefill_c_kernel, params); - } - if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerWriteBlock)) { - LaunchKernel(num_w_blocks, kBlockSize_W, device) // - .enable_pdl(kUsePDL)(prefill_w_kernel, params); - } - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh deleted file mode 100644 index b497470606..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online.cuh +++ /dev/null @@ -1,726 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include -#include - -#include -#include -#include - -namespace device::compress { - -/// \brief Plan entry for online compress 128 prefill. -/// Each entry describes a contiguous segment of tokens that lies inside a -/// single 128-chunk. Multiple segments can map to the same batch id when the -/// extend tokens span chunk boundaries. -/// -/// **Layout compatibility:** the field order/types match `PrefillPlan` so that -/// downstream kernels (e.g. `fused_norm_rope` in `CompressExtend` mode) can -/// consume the compress_plan tensor as-if it were a `PrefillPlan` tensor -- -/// they only read `ragged_id` and `position`, both of which carry identical -/// semantics here (the LAST token of the segment in q-ragged and global -/// coordinates respectively). -/// -/// Note that `window_len` here means "number of real tokens in this segment" -/// (1..128), which differs from `PrefillPlan::window_len`. Downstream kernels -/// that share the tensor MUST NOT read it under that name. -struct alignas(16) OnlinePrefillPlan { - /// \brief Ragged-q position of the LAST token in this segment. - /// Equal to `segment_start_ragged + window_len - 1`. - uint32_t ragged_id; - /// \brief Index into the `indices` / `load_indices` arrays. - uint32_t batch_id; - /// \brief Global position of the LAST token in this segment. - /// For compress plans, `position % 128 == 127` (chunk-closing); for write - /// plans, `position % 128 < 127`. - uint32_t position; - /// \brief Number of real tokens in this segment (1..128). - /// The first segment token sits at `position - window_len + 1` (global) and - /// at `ragged_id - window_len + 1` (ragged). - uint32_t window_len; -}; - -static_assert(alignof(OnlinePrefillPlan) == alignof(PrefillPlan)); -static_assert(sizeof(OnlinePrefillPlan) == sizeof(PrefillPlan)); - -} // namespace device::compress - -namespace host::compress { - -using device::compress::OnlinePrefillPlan; -using OnlinePrefillPlanTensorDtype = uint8_t; -inline constexpr int64_t kOnlinePrefillPlanDim = 16; - -static_assert(alignof(OnlinePrefillPlan) == sizeof(OnlinePrefillPlan)); -static_assert(sizeof(OnlinePrefillPlan) == kOnlinePrefillPlanDim * sizeof(OnlinePrefillPlanTensorDtype)); - -} // namespace host::compress - -namespace { - -using OnlinePlan = device::compress::OnlinePrefillPlan; -using IndiceT = int32_t; - -/// \brief Need to reduce register usage to increase occupancy -struct Compress128OnlineDecodeParams { - /** \brief Shape: `[num_indices, 1, head_dim * 3 (max, sum, kv) ]` \n */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[batch_size, head_dim * 2]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[batch_size, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[128, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, ]` */ - const IndiceT* __restrict__ seq_lens; - /** \NOTE: `batch_size` <= `num_indices` */ - uint32_t batch_size; -}; - -/// \brief Need to reduce register usage to increase occupancy -struct Compress128OnlinePrefillParams { - /** \brief Shape: `[num_indices, 1, head_dim * 3 (max, sum, kv) ]` \n */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[num_q_tokens, head_dim * 2]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[num_q_tokens, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[128, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ load_indices; - /// \brief Plan for segments that close a chunk (write to `kv_compressed_output`). - /// Shape: `[num_compress, 16]` (uint8). - const OnlinePlan* __restrict__ compress_plan; - /// \brief Plan for the trailing partial segment of each batch (write back to - /// `kv_score_buffer`). Shape: `[num_write, 16]` (uint8). - const OnlinePlan* __restrict__ write_plan; - uint32_t num_compress; - uint32_t num_write; -}; - -// 4 elements per thread, kHeadDim / 4 threads per block -template -__global__ void flash_c128_online_decode(const __grid_constant__ Compress128OnlineDecodeParams params) { - using namespace device; - constexpr uint32_t kVecSize = 4; - constexpr uint32_t kBlockSize = kHeadDim / kVecSize; - using Vec = AlignedVector; - const auto gmem = tile::Memory::cta(kBlockSize); - const auto batch_id = blockIdx.x; - const auto index = params.indices[batch_id]; - const auto seq_len = params.seq_lens[batch_id]; - - const auto kv_score_buffer = static_cast(params.kv_score_buffer); - const auto kv_buf = kv_score_buffer + index * (kHeadDim * 3); - const auto kv_score_input = static_cast(params.kv_score_input); - const auto kv_src = kv_score_input + batch_id * (kHeadDim * 2); - - /// NOTE: kv_score_buffer layout is [max, sum, kv] (slot 0 / 1 / 2). Reads, - /// writes, and the prefill kernel must all agree on this order. - const auto max_score_vec = gmem.load(kv_buf, 0); - const auto sum_score_vec = gmem.load(kv_buf, 1); - const auto old_kv_vec = gmem.load(kv_buf, 2); - - /// NOTE: kv_score_input layout is | kv | score | (head_dim each), matching - /// the offline c128 kernel and the online prefill kernel. - const auto new_kv_vec = gmem.load(kv_src, 0); - const auto new_score_raw_vec = gmem.load(kv_src, 1); - - /// NOTE: the new token sits at global position `seq_len - 1`, so its - /// position inside the 128-chunk is `(seq_len - 1) % 128`. The previous - /// `seq_len % 128` was off by one (`bias[127]` vs `bias[0]`, etc.). - const auto pos_in_chunk = (seq_len - 1) % 128; - const auto bias_vec = gmem.load(params.score_bias, pos_in_chunk); - - Vec out_kv_vec; - Vec out_max_vec; - Vec out_sum_vec; - if (pos_in_chunk != 0) { - // Mid-chunk: combine prior partial state with the new token via online softmax. -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const auto old_max = max_score_vec[i]; - const auto old_kv = old_kv_vec[i]; - const auto new_score = new_score_raw_vec[i] + bias_vec[i]; - const auto new_kv = new_kv_vec[i]; - const auto new_max = fmax(old_max, new_score); - const auto old_sum = sum_score_vec[i] * expf(old_max - new_max); - const auto new_exp = expf(new_score - new_max); - const auto new_sum = old_sum + new_exp; - out_kv_vec[i] = (old_kv * old_sum + new_kv * new_exp) / new_sum; - out_max_vec[i] = new_max; - out_sum_vec[i] = new_sum; - } - } else { - // First token of a new 128-chunk: initialize state with this token alone. -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - out_kv_vec[i] = new_kv_vec[i]; - out_max_vec[i] = new_score_raw_vec[i] + bias_vec[i]; - out_sum_vec[i] = 1.0f; // exp(score - max) with max == score - } - } - - if (pos_in_chunk == 127) { - // Chunk just closed: emit the compressed kv. No need to update the buffer - // -- the next chunk's first token will overwrite it. - const auto kv_out = static_cast(params.kv_compressed_output) + batch_id * kHeadDim; - gmem.store(kv_out, out_kv_vec); - } else { - // Otherwise persist the running [max, sum, kv] state for the next step. - gmem.store(kv_buf, out_max_vec, 0); - gmem.store(kv_buf, out_sum_vec, 1); - gmem.store(kv_buf, out_kv_vec, 2); - } -} - -constexpr int32_t kTileElements = 2; // split (along head-dim) -/// \brief Each warp will handle this many elements (split along softmax-128) -constexpr int32_t kElementsPerWarp = 8; -constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; -constexpr uint32_t kPrefillBlockSize = device::kWarpThreads * kNumWarps; -using PrefillStorage = device::AlignedVector; - -struct Compress128SharedBuffer { - using Storage = device::AlignedVector; - Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict - SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { - return data[warp_id][lane_id]; - } - SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { - return data[warp_id][lane_id][tile_id]; - } -}; - -template -SGL_DEVICE void c128_prefill_forward( - const PrefillStorage (&kv)[kElementsPerWarp], - const PrefillStorage (&score)[kElementsPerWarp], - float* kv_out, - float* max_out, - float* sum_out, - const uint32_t warp_id, - const uint32_t lane_id) { - using namespace device; - - /// NOTE: part 2: safe online softmax + weighted sum - using TmpStorage = typename Compress128SharedBuffer::Storage; - __shared__ Compress128SharedBuffer s_local_val_max; - __shared__ Compress128SharedBuffer s_local_exp_sum; - __shared__ Compress128SharedBuffer s_local_product; - - TmpStorage tmp_val_max; - TmpStorage tmp_exp_sum; - TmpStorage tmp_product; - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - float score_fp32[kElementsPerWarp]; - -#pragma unroll - for (int32_t j = 0; j < kElementsPerWarp; ++j) { - score_fp32[j] = score[j][i]; - } - - float max_value = score_fp32[0]; - float sum_exp_value = 0.0f; - -#pragma unroll - for (int32_t j = 1; j < kElementsPerWarp; ++j) { - const auto fp32_score = score_fp32[j]; - max_value = fmaxf(max_value, fp32_score); - } - - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - const auto fp32_score = score_fp32[j]; - const auto exp_score = expf(fp32_score - max_value); - sum_product += cast(kv[j][i]) * exp_score; - sum_exp_value += exp_score; - } - - tmp_val_max[i] = max_value; - tmp_exp_sum[i] = sum_exp_value; - tmp_product[i] = sum_product; - } - - // naturally aligned, so no bank conflict - s_local_val_max(warp_id, lane_id) = tmp_val_max; - s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; - s_local_product(warp_id, lane_id) = tmp_product; - - __syncthreads(); - - /// NOTE: part 3: online softmax - /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce - /// each reduce will consume `kNumWarps` threads (use partial warp reduction) - constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; - constexpr uint32_t kIteration = kReductionCount / kPrefillBlockSize; - -#pragma unroll - for (uint32_t i = 0; i < kIteration; ++i) { - /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` - const uint32_t j = i * kPrefillBlockSize + warp_id * kWarpThreads + lane_id; - /// NOTE: Range `[0, kNumWarps)` - const uint32_t local_warp_id = j % kNumWarps; - /// NOTE: Range `[0, kTileElements * kWarpThreads)` - const uint32_t local_elem_id = j / kNumWarps; - /// NOTE: Range `[0, kTileElements)` - const uint32_t local_tile_id = local_elem_id % kTileElements; - /// NOTE: Range `[0, kWarpThreads)` - const uint32_t local_lane_id = local_elem_id / kTileElements; - /// NOTE: each warp will access the whole tile (all `kTileElements`) - /// and for different lanes, the memory access only differ in `local_warp_id` - /// so there's no bank conflict in shared memory access. - static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); - const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); - const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); - const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); - const auto global_val_max = warp::reduce_max(local_val_max); - const auto rescale = expf(local_val_max - global_val_max); - const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); - const auto final_scale = rescale / global_exp_sum; - const auto global_product = warp::reduce_sum(local_product * final_scale); - kv_out[local_elem_id] = global_product; - if constexpr (kNeedData) { - max_out[local_elem_id] = global_val_max; - sum_out[local_elem_id] = global_exp_sum; - } - } - if constexpr (kNeedData) __syncthreads(); -} - -/// \brief Sentinel score for padded positions in a 128-segment. -/// Must be finite so that `score - max` never produces NaN even when an -/// entire warp has only padded positions. -constexpr float kPadScore = -FLT_MAX; - -/// \brief Online compress 128 prefill. Two passes share this body: -/// - `kWrite=false` (compress pass): handles segments that close a chunk. -/// May load prior partial state from the buffer, but never writes to it, -/// so concurrent blocks can read the same slot without racing. -/// - `kWrite=true` (write pass): handles the trailing partial segment of each -/// batch. Each batch contributes at most one such plan, so concurrent blocks -/// touch disjoint buffer slots. -/// -/// The two passes MUST run as separate kernel launches (in stream order) so -/// that all reads in pass 1 finish before any writes in pass 2 start. -template -__global__ __launch_bounds__(kPrefillBlockSize, 2) // - void flash_c128_online_prefill(const __grid_constant__ Compress128OnlinePrefillParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); - - /// NOTE: the compiler folds the if-else at compile time. - const auto num_plans = kWrite ? params.num_write : params.num_compress; - const auto plan_ptr = kWrite ? params.write_plan : params.compress_plan; - const uint32_t global_id = blockIdx.x; - const uint32_t global_pid = global_id / kNumSplit; // plan id - const uint32_t global_sid = global_id % kNumSplit; // split id - if (global_pid >= num_plans) return; - const auto [ragged_id, batch_id, position, window_len] = plan_ptr[global_pid]; - if (ragged_id == 0xFFFFFFFFu) [[unlikely]] - return; - - const uint32_t warp_id = threadIdx.x / kWarpThreads; - const uint32_t lane_id = threadIdx.x % kWarpThreads; - const int32_t split_offset = global_sid * kTileDim; // int32 is enough - - const auto kv_score_buffer = static_cast(params.kv_score_buffer); - const auto kv_score_input = static_cast(params.kv_score_input); - const auto kv_compressed_output = static_cast(params.kv_compressed_output); - const auto score_bias_base = static_cast(params.score_bias); - - constexpr int64_t kElementSize = kHeadDim * 2; // | kv | score | - const uint32_t chunk_offset = (position % 128u) + 1u - window_len; - const uint32_t window_end = chunk_offset + window_len; // exclusive, in [1, 128] - const int32_t segment_start = ragged_id - (position % 128u); // can be negative, but safe - const int32_t load_index = chunk_offset != 0 ? params.load_indices[batch_id] : -1; - const int32_t store_index = kWrite ? params.indices[batch_id] : -1; - - PDLWaitPrimary(); - - // 2 * 8 = 16 register per elem. in theory we should consume 48 register here - PrefillStorage kv[kElementsPerWarp]; - PrefillStorage score[kElementsPerWarp]; - PrefillStorage bias[kElementsPerWarp]; - const auto warp_offset = warp_id * kElementsPerWarp; - -#pragma unroll - for (uint32_t i = 0; i < kElementsPerWarp; ++i) { - const uint32_t j = i + warp_offset; - if (j >= chunk_offset && j < window_end) { - const auto kv_src_ptr = kv_score_input + (segment_start + j) * kElementSize + split_offset; - const auto score_src_ptr = kv_src_ptr + kHeadDim; - const auto bias_src_ptr = score_bias_base + j * kHeadDim + split_offset; - kv[i].load(kv_src_ptr, lane_id); - score[i].load(score_src_ptr, lane_id); - bias[i].load(bias_src_ptr, lane_id); - } - } - -#pragma unroll - for (uint32_t i = 0; i < kElementsPerWarp; ++i) { - const uint32_t j = i + warp_offset; - const bool is_valid = (j >= chunk_offset && j < window_end); -#pragma unroll - for (uint32_t ii = 0; ii < kTileElements; ++ii) { - score[i][ii] = is_valid ? score[i][ii] + bias[i][ii] : kPadScore; - /// NOTE: must zero out kv on padded slots -- `c128_prefill_forward` - /// computes `kv * exp_score` where `exp_score = expf(-FLT_MAX - max) ??? 0`, - /// and IEEE-754 makes `NaN * 0 = NaN` / `+-inf * 0 = NaN`. An - /// uninitialized register can hold a NaN/inf bit pattern, so without - /// this reset a single padded warp can poison the whole softmax. - kv[i][ii] = is_valid ? kv[i][ii] : 0.0f; - } - } - - __shared__ alignas(16) float seg_kv[kTileDim]; - __shared__ alignas(16) float seg_max[kTileDim]; - __shared__ alignas(16) float seg_sum[kTileDim]; - - c128_prefill_forward(kv, score, seg_kv, seg_max, seg_sum, warp_id, lane_id); - - PDLTriggerSecondary(); - - if (warp_id == 0) { - PrefillStorage out_kv_vec, out_max_vec, out_sum_vec; - out_kv_vec.load(seg_kv, lane_id); - out_max_vec.load(seg_max, lane_id); - out_sum_vec.load(seg_sum, lane_id); - if (chunk_offset != 0) { - /// NOTE: load (max, sum, kv) of the in-progress chunk for this index. - /// `load_indices` may differ from `indices` when the prior partial state - /// lives on a different slot than the slot we ultimately write to. - const auto buf_load = kv_score_buffer + load_index * (kHeadDim * 3) + split_offset; - PrefillStorage buf_max_vec, buf_sum_vec, buf_kv_vec; - buf_max_vec.load(buf_load + 0 * kHeadDim, lane_id); - buf_sum_vec.load(buf_load + 1 * kHeadDim, lane_id); - buf_kv_vec.load(buf_load + 2 * kHeadDim, lane_id); -#pragma unroll - for (uint32_t ii = 0; ii < kTileElements; ++ii) { - const float m1 = buf_max_vec[ii]; - const float s1 = buf_sum_vec[ii]; - const float k1 = buf_kv_vec[ii]; - const float m2 = out_max_vec[ii]; - const float s2 = out_sum_vec[ii]; - const float k2 = out_kv_vec[ii]; - const float new_max = fmaxf(m1, m2); - const float new_s1 = s1 * expf(m1 - new_max); - const float new_s2 = s2 * expf(m2 - new_max); - const float new_sum = new_s1 + new_s2; - const float new_kv = (k1 * new_s1 + k2 * new_s2) / new_sum; - out_max_vec[ii] = new_max; - out_sum_vec[ii] = new_sum; - out_kv_vec[ii] = new_kv; - } - } - - if constexpr (kWrite) { - const auto buf_store = kv_score_buffer + store_index * (kHeadDim * 3) + split_offset; - reinterpret_cast(buf_store + 0 * kHeadDim)[lane_id] = out_max_vec; - reinterpret_cast(buf_store + 1 * kHeadDim)[lane_id] = out_sum_vec; - reinterpret_cast(buf_store + 2 * kHeadDim)[lane_id] = out_kv_vec; - } else { - const auto out_ptr = kv_compressed_output + ragged_id * kHeadDim + split_offset; - reinterpret_cast(out_ptr)[lane_id] = out_kv_vec; - } - } -} - -template -struct FlashCompress128OnlineKernel { - static constexpr auto decode_kernel = flash_c128_online_decode; - template - static constexpr auto prefill_kernel = flash_c128_online_prefill; - static constexpr auto prefill_c_kernel = prefill_kernel; - static constexpr auto prefill_w_kernel = prefill_kernel; - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static constexpr uint32_t kDecodeBlockSize = kHeadDim / 4; - - static void run_decode( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::Optional /* UNUSED */) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) - .with_dtype() - .with_device(device) - .verify(kv_score_buffer); - TensorMatcher({B, kHeadDim * 2}) // kv score input - .with_dtype() - .with_device(device) - .verify(kv_score_input); - TensorMatcher({B, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device) - .verify(ape); - TensorMatcher({B}).with_dtype().with_device(device).verify(indices); - TensorMatcher({B}).with_dtype().with_device(device).verify(seq_lens); - - const auto batch_size = static_cast(B.unwrap()); - const auto params = Compress128OnlineDecodeParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .batch_size = batch_size, - }; - LaunchKernel(batch_size, kDecodeBlockSize, device.unwrap()) // - .enable_pdl(kUsePDL)(decode_kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView compress_plan, - const tvm::ffi::TensorView write_plan, - const tvm::ffi::Optional extra) { - using namespace host; - using host::compress::kOnlinePrefillPlanDim; - using host::compress::OnlinePrefillPlanTensorDtype; - - auto B = SymbolicSize{"batch_size"}; - auto N = SymbolicSize{"num_q_tokens"}; - auto X = SymbolicSize{"compress_tokens"}; - auto Y = SymbolicSize{"write_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) ??? 2D - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({N, kHeadDim * 2}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({N, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - TensorMatcher({B}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - TensorMatcher({X, kOnlinePrefillPlanDim}) // compress plan - .with_dtype() - .with_device(device_) - .verify(compress_plan); - TensorMatcher({Y, kOnlinePrefillPlanDim}) // write plan - .with_dtype() - .with_device(device_) - .verify(write_plan); - - /// NOTE: `extra` is `load_indices`. When the previous partial state lives - /// on a slot different from the destination slot (e.g. paged buffers), the - /// caller must supply this; otherwise it defaults to `indices`. - const auto load_indices = extra.value_or(indices); - TensorMatcher({B}).with_dtype().with_device(device_).verify(load_indices); - - const auto device = device_.unwrap(); - const auto num_c = static_cast(X.unwrap()); - const auto num_w = static_cast(Y.unwrap()); - const auto params = Compress128OnlinePrefillParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .load_indices = static_cast(load_indices.data_ptr()), - .compress_plan = static_cast(compress_plan.data_ptr()), - .write_plan = static_cast(write_plan.data_ptr()), - .num_compress = num_c, - .num_write = num_w, - }; - - /// NOTE: pass 1 reads the buffer (for the first segment of each batch - /// that started mid-chunk) and writes only to `kv_compressed_output`. - /// Pass 2 then writes the trailing partial state of each batch back to - /// the buffer. Stream serialization between the two launches enforces - /// read-before-write on shared buffer slots. - if (const auto num_c_blocks = num_c * kNumSplit) { - LaunchKernel(num_c_blocks, kPrefillBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_c_kernel, params); - } - if (const auto num_w_blocks = num_w * kNumSplit) { - LaunchKernel(num_w_blocks, kPrefillBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_w_kernel, params); - } - } -}; - -} // namespace - -namespace host::compress { - -using OnlinePlanResult = tvm::ffi::Tuple; - -struct OnlinePrefillCompressParams { - OnlinePrefillPlan* __restrict__ compress_plan; - OnlinePrefillPlan* __restrict__ write_plan; - const int64_t* __restrict__ seq_lens; - const int64_t* __restrict__ extend_lens; - uint32_t batch_size; - uint32_t num_tokens; -}; - -/// \brief Build the compress + write plans for online compress 128 prefill. -/// -/// Each batch's `[prefix_len, prefix_len + extend_len)` range is split at -/// 128-aligned boundaries. Every resulting segment falls into one of: -/// - **compress**: closes a 128-chunk (`chunk_offset + window_len == 128`). -/// These plans only read the buffer (when starting mid-chunk) and write the -/// compressed kv to `kv_compressed_output`. -/// - **write**: trailing partial of the batch (`chunk_offset + window_len < 128`). -/// May read the buffer and always writes the new partial state back to it. -/// Each batch produces at most one such plan. -/// -/// The two plans MUST be dispatched as separate kernel launches in stream -/// order so that pass-1 reads of a buffer slot complete before any pass-2 -/// write of the same slot. -inline OnlinePlanResult plan_online_prefill_host(const OnlinePrefillCompressParams& params, const bool use_cuda_graph) { - const auto& [compress_plan, write_plan, seq_lens, extend_lens, batch_size, num_tokens] = params; - - uint32_t counter = 0; - uint32_t compress_count = 0; - uint32_t write_count = 0; - for (const auto i : irange(batch_size)) { - const uint32_t seq_len = static_cast(seq_lens[i]); - const uint32_t extend_len = static_cast(extend_lens[i]); - RuntimeCheck(0 < extend_len && extend_len <= seq_len); - const uint32_t prefix_len = seq_len - extend_len; - const uint32_t end_pos = prefix_len + extend_len; - /// NOTE: split the extend range into per-128-chunk segments. Each segment - /// stays inside one chunk, so the kernel can decide load/store from - /// `chunk_offset` and `window_len` alone. - uint32_t pos = prefix_len; - while (pos < end_pos) { - const uint32_t chunk_start = (pos / 128u) * 128u; - const uint32_t seg_end = std::min(end_pos, chunk_start + 128u); // exclusive - const uint32_t seg_len = seg_end - pos; - const uint32_t chunk_off = pos - chunk_start; - /// NOTE: store last-token coordinates so that downstream consumers - /// (e.g. `fused_norm_rope`) can read `ragged_id` and `position` with the - /// same semantics as `PrefillPlan`. The segment start is recoverable as - /// `ragged_id - window_len + 1` and `position - window_len + 1`. - const uint32_t last_pos = seg_end - 1; - const uint32_t last_ragged = counter + (last_pos - prefix_len); - const auto plan = OnlinePrefillPlan{ - .ragged_id = last_ragged, - .batch_id = i, - .position = last_pos, - .window_len = seg_len, - }; - if (chunk_off + seg_len == 128u) { - // full chunk, must be complete, maybe read the buffer, no write - RuntimeCheck(compress_count < num_tokens); - compress_plan[compress_count++] = plan; - } else { - // last chunk, must be incomplete, maybe read the buffer, must write - RuntimeCheck(write_count < num_tokens); - write_plan[write_count++] = plan; - } - pos = seg_end; - } - counter += extend_len; - } - RuntimeCheck(counter == num_tokens, "input size ", counter, " != num_q_tokens ", num_tokens); - if (!use_cuda_graph) return OnlinePlanResult{compress_count, write_count}; - /// NOTE: pad both plans with sentinel entries so cuda-graph runs always see - /// the same number of blocks. The kernel skips plans whose `ragged_id` is -1. - constexpr auto kInvalid = static_cast(-1); - constexpr auto kInvalidPlan = OnlinePrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; - for (const auto i : irange(compress_count, num_tokens)) { - compress_plan[i] = kInvalidPlan; - } - for (const auto i : irange(write_count, num_tokens)) { - write_plan[i] = kInvalidPlan; - } - return OnlinePlanResult{num_tokens, num_tokens}; -} - -inline OnlinePlanResult plan_online_prefill( - const tvm::ffi::TensorView extend_lens, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView compress_plan, - const tvm::ffi::TensorView write_plan, - const bool use_cuda_graph) { - auto N = SymbolicSize{"batch_size"}; - auto M = SymbolicSize{"num_tokens"}; - auto device = SymbolicDevice{}; - /// NOTE: only host (CPU/cuda-host) planning is implemented for now. The - device.set_options(); - TensorMatcher({N}) // - .with_dtype() - .with_device(device) - .verify(extend_lens) - .verify(seq_lens); - TensorMatcher({M, kOnlinePrefillPlanDim}) // - .with_dtype() - .with_device(device) - .verify(compress_plan) - .verify(write_plan); - const auto params = OnlinePrefillCompressParams{ - .compress_plan = static_cast(compress_plan.data_ptr()), - .write_plan = static_cast(write_plan.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .extend_lens = static_cast(extend_lens.data_ptr()), - .batch_size = static_cast(N.unwrap()), - .num_tokens = static_cast(M.unwrap()), - }; - return plan_online_prefill_host(params, use_cuda_graph); -} - -} // namespace host::compress - -namespace { - -[[maybe_unused]] -constexpr auto& plan_compress_online_prefill = host::compress::plan_online_prefill; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh deleted file mode 100644 index 71e600dc39..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_online_v2.cuh +++ /dev/null @@ -1,875 +0,0 @@ -#include -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace { - -using PlanD = device::compress::DecodePlan; -using PlanC = device::compress::CompressPlan; - -// --------------------------------------------------------------------------- -// Decode kernel: 1 token / batch. Each block handles one batch. -// 4 elements per thread -> kBlockSize = head_dim / 4. -// --------------------------------------------------------------------------- - -struct Compress128OnlineDecodeParams { - void* __restrict__ kv_score_buffer; // [num_slots, 1, head_dim * 3] - const void* __restrict__ kv_score_input; // [batch_size, head_dim * 2] - void* __restrict__ kv_compressed_output; // [batch_size, head_dim] - const void* __restrict__ score_bias; // [128, head_dim] - const PlanD* __restrict__ plan_d; - uint32_t batch_size; -}; - -template -__global__ void flash_c128_online_decode_v2(const __grid_constant__ Compress128OnlineDecodeParams params) { - using namespace device; - constexpr uint32_t kVecSize = 4; - constexpr uint32_t kBlockSize = kHeadDim / kVecSize; - using Vec = AlignedVector; - const auto gmem = tile::Memory::cta(kBlockSize); - const auto batch_id = blockIdx.x; - if (batch_id >= params.batch_size) return; - - // Wait for the plan-finalize kernel to publish `plan.read_page_0 / write_loc` - // before reading the plan. The plan kernel runs on the same stream and does - // NOT issue a PDL trigger, so launching this kernel with PDL means our - // pre-wait global reads can race with the plan kernel's writes. - PDLWaitPrimary(); - - const auto plan = params.plan_d[batch_id]; - const auto pos_in_chunk = (plan.seq_len - 1) % 128; - - const auto kv_score_buffer = static_cast(params.kv_score_buffer); - const auto kv_score_input = static_cast(params.kv_score_input); - const auto kv_load_buf = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3); - const auto kv_store_buf = kv_score_buffer + plan.write_loc * (kHeadDim * 3); - const auto kv_src = kv_score_input + batch_id * (kHeadDim * 2); - - // Buffer layout: [max | sum | kv] (slot 0 / 1 / 2 of the head_dim*3 row). - const auto new_kv_vec = gmem.load(kv_src, 0); - const auto new_score_raw_vec = gmem.load(kv_src, 1); - const auto bias_vec = gmem.load(params.score_bias, pos_in_chunk); - - Vec out_kv_vec; - Vec out_max_vec; - Vec out_sum_vec; - if (pos_in_chunk != 0) { - // Mid-chunk: combine prior partial state with the new token. - const auto max_score_vec = gmem.load(kv_load_buf, 0); - const auto sum_score_vec = gmem.load(kv_load_buf, 1); - const auto old_kv_vec = gmem.load(kv_load_buf, 2); -#pragma unroll - for (uint32_t i = 0; i < kVecSize; ++i) { - const auto old_max = max_score_vec[i]; - const auto old_kv = old_kv_vec[i]; - const auto new_score = new_score_raw_vec[i] + bias_vec[i]; - const auto new_kv = new_kv_vec[i]; - const auto new_max = fmaxf(old_max, new_score); - const auto old_sum = sum_score_vec[i] * expf(old_max - new_max); - const auto new_exp = expf(new_score - new_max); - const auto new_sum = old_sum + new_exp; - out_kv_vec[i] = (old_kv * old_sum + new_kv * new_exp) / new_sum; - out_max_vec[i] = new_max; - out_sum_vec[i] = new_sum; - } - } else { - // First token of a new chunk: state == this token alone. -#pragma unroll - for (uint32_t i = 0; i < kVecSize; ++i) { - out_kv_vec[i] = new_kv_vec[i]; - out_max_vec[i] = new_score_raw_vec[i] + bias_vec[i]; - out_sum_vec[i] = 1.0f; - } - } - - if (pos_in_chunk == 127) { - // Chunk just closed: emit compressed kv, no buffer update. - const auto kv_out = static_cast(params.kv_compressed_output) + batch_id * kHeadDim; - gmem.store(kv_out, out_kv_vec); - } else { - gmem.store(kv_store_buf, out_max_vec, 0); - gmem.store(kv_store_buf, out_sum_vec, 1); - gmem.store(kv_store_buf, out_kv_vec, 2); - } -} - -// --------------------------------------------------------------------------- -// Prefill kernel: 1 segment / block. Two passes (compress + write) share the -// kernel template, parameterized by `kWrite`. -// 16 warps per block; each warp handles 8 of the 128 chunk positions. -// --------------------------------------------------------------------------- - -constexpr int32_t kTileElements = 2; // split along head-dim -constexpr int32_t kElementsPerWarp = 8; // split along the 128-chunk -constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; -constexpr uint32_t kPrefillBlockSize = device::kWarpThreads * kNumWarps; -using PrefillStorage = device::AlignedVector; - -struct Compress128OnlinePrefillParams { - void* __restrict__ kv_score_buffer; // [num_slots, 1, head_dim * 3] - const void* __restrict__ kv_score_input; // [num_q_tokens, head_dim * 2] - void* __restrict__ kv_compressed_output; // [num_compress, head_dim] - const void* __restrict__ score_bias; // [128, head_dim] - const PlanC* __restrict__ plan_c; // close-chunk segments - const PlanC* __restrict__ plan_w; // trailing partial segments - uint32_t num_compress; - uint32_t num_write; -}; - -struct Compress128SharedBuffer { - using Storage = device::AlignedVector; - Storage data[kNumWarps][device::kWarpThreads + 1]; // +1 to avoid bank conflict - SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { - return data[warp_id][lane_id]; - } - SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { - return data[warp_id][lane_id][tile_id]; - } -}; - -/// \brief Sentinel score for padded positions in a 128-segment. -constexpr float kPadScore = -FLT_MAX; - -[[maybe_unused]] -SGL_DEVICE void c128_prefill_segment_softmax( - const PrefillStorage (&kv)[kElementsPerWarp], - const PrefillStorage (&score)[kElementsPerWarp], - float* seg_kv, - float* seg_max, - float* seg_sum, - const uint32_t warp_id, - const uint32_t lane_id) { - using namespace device; - - // Per-warp running state (max, sum, kv) for kTileElements head-dim slots. - using TmpStorage = typename Compress128SharedBuffer::Storage; - __shared__ Compress128SharedBuffer s_local_val_max; - __shared__ Compress128SharedBuffer s_local_exp_sum; - __shared__ Compress128SharedBuffer s_local_product; - - TmpStorage tmp_val_max; - TmpStorage tmp_exp_sum; - TmpStorage tmp_product; - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - float score_fp32[kElementsPerWarp]; -#pragma unroll - for (int32_t j = 0; j < kElementsPerWarp; ++j) { - score_fp32[j] = score[j][i]; - } - float max_value = score_fp32[0]; -#pragma unroll - for (int32_t j = 1; j < kElementsPerWarp; ++j) { - max_value = fmaxf(max_value, score_fp32[j]); - } - float sum_exp_value = 0.0f; - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < kElementsPerWarp; ++j) { - const auto exp_score = expf(score_fp32[j] - max_value); - sum_product += kv[j][i] * exp_score; - sum_exp_value += exp_score; - } - tmp_val_max[i] = max_value; - tmp_exp_sum[i] = sum_exp_value; - tmp_product[i] = sum_product; - } - - // Aligned writes (no bank conflict thanks to `+1` padding). - s_local_val_max(warp_id, lane_id) = tmp_val_max; - s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; - s_local_product(warp_id, lane_id) = tmp_product; - - __syncthreads(); - - // Cross-warp reduction. Same recipe as c128_online.cuh: each block-thread - // pair reduces a (tile_id, lane_id) slot using a kNumWarps-wide warp shuffle. - constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; - constexpr uint32_t kIteration = kReductionCount / kPrefillBlockSize; - static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); - -#pragma unroll - for (uint32_t i = 0; i < kIteration; ++i) { - const uint32_t j = i * kPrefillBlockSize + warp_id * kWarpThreads + lane_id; - const uint32_t local_warp_id = j % kNumWarps; - const uint32_t local_elem_id = j / kNumWarps; - const uint32_t local_tile_id = local_elem_id % kTileElements; - const uint32_t local_lane_id = local_elem_id / kTileElements; - const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); - const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); - const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); - const auto global_val_max = warp::reduce_max(local_val_max); - const auto rescale = expf(local_val_max - global_val_max); - const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); - const auto final_scale = rescale / global_exp_sum; - const auto global_product = warp::reduce_sum(local_product * final_scale); - seg_kv[local_elem_id] = global_product; - seg_max[local_elem_id] = global_val_max; - seg_sum[local_elem_id] = global_exp_sum; - } - __syncthreads(); -} - -/// \brief Online compress 128 prefill v2. -/// -/// `kWrite=false` (compress pass): handles segments that close a 128-chunk. -/// Reads optional prior state from `read_page_0` (-1 = none), emits compressed -/// kv to `kv_compressed_output[plan_id]` (compact). -/// `kWrite=true` (write pass) : handles trailing partial segments. -/// Reads optional prior state from `read_page_0` (-1 = none), writes new -/// running state to `read_page_1`. -template -__global__ __launch_bounds__(kPrefillBlockSize, 2) // - void flash_c128_online_prefill_v2(const __grid_constant__ Compress128OnlinePrefillParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 64 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static_assert(kHeadDim % kTileDim == 0); - - // Compile-time fold to the right plan list. - const auto num_plans = kWrite ? params.num_write : params.num_compress; - const auto plan_ptr = kWrite ? params.plan_w : params.plan_c; - const uint32_t global_id = blockIdx.x; - const uint32_t global_pid = global_id / kNumSplit; - const uint32_t global_sid = global_id % kNumSplit; - if (global_pid >= num_plans) return; - - const uint32_t warp_id = threadIdx.x / kWarpThreads; - const uint32_t lane_id = threadIdx.x % kWarpThreads; - const int32_t split_offset = global_sid * kTileDim; - - // The previous kernel (plan-finalize stage 1) does NOT issue a PDL trigger, - // so PDLWaitPrimary effectively waits for stage 1 to complete. Read the plan - // AFTER the wait so the freshly-written `read_page_0` (= state-pool slot) is - // visible. Reading it before the wait is a real race -- with PDL enabled the - // kernel can begin executing before stage 1's stores propagate, and we'd see - // the stage-0 batch_id placeholder in `read_page_0` instead of the slot. - PDLWaitPrimary(); - - const auto plan = plan_ptr[global_pid]; - if (plan.is_invalid()) [[unlikely]] - return; - - const auto kv_score_buffer = static_cast(params.kv_score_buffer); - const auto kv_score_input = static_cast(params.kv_score_input); - const auto kv_compressed_output = static_cast(params.kv_compressed_output); - const auto score_bias_base = static_cast(params.score_bias); - - constexpr int64_t kElementSize = kHeadDim * 2; // | kv | score | - - // The plan stores last-token coordinates; segment start is recoverable as - // ragged_id - window_len + 1. - const uint32_t window_len = plan.buffer_len; - const uint32_t position = plan.seq_len - 1; - const uint32_t pos_in_chunk_end = (position % 128u) + 1u; // exclusive, in [1, 128] - const uint32_t chunk_offset = pos_in_chunk_end - window_len; // in [0, 127] - const int32_t segment_start_ragged = static_cast(plan.ragged_id) - static_cast(position % 128u); - - // --- Stage 1: load kv / score / bias for this warp's 8 chunk positions. - PrefillStorage kv[kElementsPerWarp]; - PrefillStorage score[kElementsPerWarp]; - PrefillStorage bias[kElementsPerWarp]; - const uint32_t warp_offset = warp_id * kElementsPerWarp; - -#pragma unroll - for (uint32_t i = 0; i < kElementsPerWarp; ++i) { - const uint32_t j = i + warp_offset; - if (j >= chunk_offset && j < pos_in_chunk_end) { - const auto kv_src_ptr = kv_score_input + (segment_start_ragged + j) * kElementSize + split_offset; - const auto score_src_ptr = kv_src_ptr + kHeadDim; - const auto bias_src_ptr = score_bias_base + j * kHeadDim + split_offset; - kv[i].load(kv_src_ptr, lane_id); - score[i].load(score_src_ptr, lane_id); - bias[i].load(bias_src_ptr, lane_id); - } - } - - // --- Stage 2: pad invalid positions. score = -FLT_MAX, kv = 0 (so that - // kv * exp(score-max) ??? 0 / 0 cleanly without producing NaN/inf). -#pragma unroll - for (uint32_t i = 0; i < kElementsPerWarp; ++i) { - const uint32_t j = i + warp_offset; - const bool is_valid = (j >= chunk_offset && j < pos_in_chunk_end); -#pragma unroll - for (uint32_t ii = 0; ii < kTileElements; ++ii) { - score[i][ii] = is_valid ? score[i][ii] + bias[i][ii] : kPadScore; - kv[i][ii] = is_valid ? kv[i][ii] : 0.0f; - } - } - - // --- Stage 3: warp-tile online softmax over the 128-position chunk. - __shared__ alignas(16) float seg_kv[kTileDim]; - __shared__ alignas(16) float seg_max[kTileDim]; - __shared__ alignas(16) float seg_sum[kTileDim]; - c128_prefill_segment_softmax(kv, score, seg_kv, seg_max, seg_sum, warp_id, lane_id); - - PDLTriggerSecondary(); - - // --- Stage 4: warp 0 folds with prior partial state (if any) and writes. - if (warp_id == 0) { - PrefillStorage out_kv_vec, out_max_vec, out_sum_vec; - out_kv_vec.load(seg_kv, lane_id); - out_max_vec.load(seg_max, lane_id); - out_sum_vec.load(seg_sum, lane_id); - - if (chunk_offset != 0 && plan.read_page_0 >= 0) { - // Combine with prior partial state for this slot. - const auto buf_load = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3) + split_offset; - PrefillStorage buf_max_vec, buf_sum_vec, buf_kv_vec; - buf_max_vec.load(buf_load + 0 * kHeadDim, lane_id); - buf_sum_vec.load(buf_load + 1 * kHeadDim, lane_id); - buf_kv_vec.load(buf_load + 2 * kHeadDim, lane_id); -#pragma unroll - for (uint32_t ii = 0; ii < kTileElements; ++ii) { - const float m1 = buf_max_vec[ii]; - const float s1 = buf_sum_vec[ii]; - const float k1 = buf_kv_vec[ii]; - const float m2 = out_max_vec[ii]; - const float s2 = out_sum_vec[ii]; - const float k2 = out_kv_vec[ii]; - const float new_max = fmaxf(m1, m2); - const float new_s1 = s1 * expf(m1 - new_max); - const float new_s2 = s2 * expf(m2 - new_max); - const float new_sum = new_s1 + new_s2; - const float new_kv = (k1 * new_s1 + k2 * new_s2) / new_sum; - out_max_vec[ii] = new_max; - out_sum_vec[ii] = new_sum; - out_kv_vec[ii] = new_kv; - } - } - - if constexpr (kWrite) { - // For trailing-partial segments the load and store slots collapse to the - // segment's own chunk slot (the request keeps a single in-progress - // chunk's running state at any time), so we reuse `read_page_0`. - const auto buf_store = kv_score_buffer + plan.read_page_0 * (kHeadDim * 3) + split_offset; - reinterpret_cast(buf_store + 0 * kHeadDim)[lane_id] = out_max_vec; - reinterpret_cast(buf_store + 1 * kHeadDim)[lane_id] = out_sum_vec; - reinterpret_cast(buf_store + 2 * kHeadDim)[lane_id] = out_kv_vec; - } else { - // Compact output: one row per compress plan, indexed by `global_pid`. - const auto out_ptr = kv_compressed_output + global_pid * kHeadDim + split_offset; - reinterpret_cast(out_ptr)[lane_id] = out_kv_vec; - } - } -} - -// --------------------------------------------------------------------------- -// Host wrapper: matches the c128_v2 / c4_v2 host API style (run_decode / -// run_prefill methods on a kernel-class template). We only expose `kHeadDim` -// + `kUsePDL`; the dtype is fixed to fp32 for the online state pool. -// --------------------------------------------------------------------------- - -template -struct FlashCompress128OnlineKernel { - static constexpr auto decode_kernel = flash_c128_online_decode_v2; - template - static constexpr auto prefill_kernel = flash_c128_online_prefill_v2; - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static constexpr uint32_t kDecodeBlockSize = kHeadDim / 4; - - static void run_decode( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_d_) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer (max, sum, kv) - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({B, kHeadDim * 2}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({B, kHeadDim}) // kv compressed output (sparse by batch_id) - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - - const auto plan_d = compress::verify_plan_d(plan_d_, B, device_); - const auto batch_size = static_cast(B.unwrap()); - if (batch_size == 0) return; - const auto params = Compress128OnlineDecodeParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_d = plan_d, - .batch_size = batch_size, - }; - LaunchKernel(batch_size, kDecodeBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(decode_kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_c_, - const tvm::ffi::TensorView plan_w_) { - using namespace host; - - auto N = SymbolicSize{"num_q_tokens"}; - auto C = SymbolicSize{"num_c_plans"}; - auto W = SymbolicSize{"num_w_plans"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 1, kHeadDim * 3}) // kv score buffer - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({N, kHeadDim * 2}) // kv score input (ragged) - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({C, kHeadDim}) // kv compressed output (compact, by plan_c index) - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - - // Both compress and write segments use PlanC layout. plan_c uses - // read_page_1=-1 (unused); plan_w uses read_page_1=store_slot. - const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); - const auto plan_w = compress::verify_plan_c(plan_w_, W, device_); - const auto device = device_.unwrap(); - const auto num_q_tokens = static_cast(N.unwrap()); - const auto num_c = static_cast(C.unwrap()); - const auto num_w = static_cast(W.unwrap()); - RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); - const auto params = Compress128OnlinePrefillParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_c = plan_c, - .plan_w = plan_w, - .num_compress = num_c, - .num_write = num_w, - }; - - // The two passes MUST be serialized in stream order: pass 1 reads slots - // that pass 2 may write to; running them in parallel would race. - if (const auto num_c_blocks = num_c * kNumSplit) { - LaunchKernel(num_c_blocks, kPrefillBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_kernel, params); - } - if (const auto num_w_blocks = num_w * kNumSplit) { - LaunchKernel(num_w_blocks, kPrefillBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_kernel, params); - } - } -}; - -} // namespace - -// =========================================================================== -// Plan builders. Mirrors the offline v2 pattern (`c_plan.cuh`): -// - Decode: a single GPU kernel reads seq_lens / req_to_token / -// req_pool_indices on device and emits the final PlanD tensor in one go. -// - Prefill: stage 0 (host, on CPU pinned memory) splits each batch's -// extend range into per-chunk segments and emits PlanC entries with the -// batch_id stashed in `read_page_0` as a placeholder. Stage 1 is a tiny -// GPU kernel that finalizes `read_page_0` to `req_to_token[rid][chunk_start]`, -// so the slot tensors never leave GPU memory. The online state pool keeps -// a single in-progress chunk per request, so each segment's load and -// store slot collapse to one value (the slot for the segment's own chunk), -// and `read_page_1` is unused. -// =========================================================================== - -namespace host::compress { - -using device::compress::CompressPlan; -using device::compress::DecodePlan; - -// --------------------------------------------------------------------------- -// Decode plan builder. -// --------------------------------------------------------------------------- - -struct OnlineDecodePlanParams { - DecodePlan* __restrict__ plan_d; - const int64_t* __restrict__ seq_lens; - const int64_t* __restrict__ req_pool_indices; - const int32_t* __restrict__ req_to_token; - const int64_t* __restrict__ full_to_swa; // (full_cache_size,) int64 - int64_t stride_r2t; - int32_t swa_page_size; - uint32_t batch_size; -}; - -__global__ void plan_c128_online_decode_kernel(const OnlineDecodePlanParams params) { - const uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= params.batch_size) return; - const auto seq_len = static_cast(params.seq_lens[idx]); - const auto rid = params.req_pool_indices[idx]; - const int32_t chunk_start = static_cast((seq_len - 1u) / 128u * 128u); - const int32_t full_loc = params.req_to_token[rid * params.stride_r2t + chunk_start]; - const int32_t swa_loc = static_cast(params.full_to_swa[full_loc]); - const int32_t slot = swa_loc / params.swa_page_size; - params.plan_d[idx] = DecodePlan{ - .seq_len = seq_len, - .write_loc = slot, - .read_page_0 = slot, - .read_page_1 = -1, - }; -} - -/// \brief Build the decode plan tensor. Caller (Python) pre-allocates -/// `plan_d_dev` as a `(batch_size, 16)` device uint8 tensor; this routine -/// only fills it. See `plan_online_prefill` for the rationale (avoid -/// `ffi::empty` + dlpack roundtrip / PyTorch caching-allocator stream -/// tracking issue that surfaces as IMA in unrelated downstream kernels). -inline void plan_online_decode( - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView req_pool_indices, - const tvm::ffi::TensorView req_to_token, - const tvm::ffi::TensorView full_to_swa, - const tvm::ffi::TensorView plan_d_dev_, - const int32_t swa_page_size) { - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - auto seq_dtype = SymbolicDType{}; - TensorMatcher({B}) // - .with_dtype(seq_dtype) - .with_device(device_) - .verify(seq_lens); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({-1, -1}) // - .with_dtype() - .with_device(device_) - .verify(req_to_token); - TensorMatcher({-1}) // - .with_dtype() - .with_device(device_) - .verify(full_to_swa); - TensorMatcher({B, sizeof(DecodePlan)}) // - .with_dtype() - .with_device(device_) - .verify(plan_d_dev_); - RuntimeCheck(swa_page_size > 0); - - const auto batch_size = static_cast(B.unwrap()); - if (batch_size == 0) return; - - const auto device = device_.unwrap(); - constexpr uint32_t kBlockSize = 256; - const uint32_t num_blocks = host::div_ceil(batch_size, kBlockSize); - const auto stride_r2t = req_to_token.stride(0); - const auto params = OnlineDecodePlanParams{ - .plan_d = static_cast(plan_d_dev_.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .req_pool_indices = static_cast(req_pool_indices.data_ptr()), - .req_to_token = static_cast(req_to_token.data_ptr()), - .full_to_swa = static_cast(full_to_swa.data_ptr()), - .stride_r2t = stride_r2t, - .swa_page_size = swa_page_size, - .batch_size = batch_size, - }; - LaunchKernel(num_blocks, kBlockSize, device)(plan_c128_online_decode_kernel, params); -} - -// --------------------------------------------------------------------------- -// Prefill plan builder: host stage 0 + GPU stage 1. -// --------------------------------------------------------------------------- - -struct OnlinePrefillStage0Params { - CompressPlan* __restrict__ plan_c; - CompressPlan* __restrict__ plan_w; - const int64_t* __restrict__ seq_lens; - const int64_t* __restrict__ extend_lens; - uint32_t batch_size; - uint32_t num_q_tokens; -}; - -inline std::tuple _plan_prefill_partial(const OnlinePrefillStage0Params& p) { - uint32_t counter = 0; - uint32_t compress_count = 0; - uint32_t write_count = 0; - for (const auto i : irange(p.batch_size)) { - const uint32_t seq_len = static_cast(p.seq_lens[i]); - const uint32_t extend_len = static_cast(p.extend_lens[i]); - RuntimeCheck(0 < extend_len && extend_len <= seq_len); - const uint32_t prefix_len = seq_len - extend_len; - const uint32_t end_pos = prefix_len + extend_len; - - uint32_t pos = prefix_len; - while (pos < end_pos) { - const uint32_t chunk_start = (pos / 128u) * 128u; - const uint32_t seg_end = std::min(end_pos, chunk_start + 128u); // exclusive - const uint32_t seg_len = seg_end - pos; - const uint32_t chunk_off = pos - chunk_start; - const uint32_t last_pos = seg_end - 1; - const uint32_t last_ragged = counter + (last_pos - prefix_len); - RuntimeCheck(last_ragged < (1u << 16), "PlanC.ragged_id is uint16; ragged ", last_ragged, " overflows"); - RuntimeCheck(seg_len <= 128u); - // Stash batch_id in `read_page_0` for stage 1 to translate. A - // chunk-aligned segment never loads, so we still need stage 1 to fill - // a slot in -- the kernel keys the load on `chunk_offset != 0`. - const auto plan = CompressPlan{ - .seq_len = last_pos + 1u, - .ragged_id = static_cast(last_ragged), - .buffer_len = static_cast(seg_len), - .read_page_0 = static_cast(i), // batch_id placeholder - .read_page_1 = -1, // unused, kept so MSB layout is stable - }; - if (chunk_off + seg_len == 128u) { - // close-chunk segment - RuntimeCheck(compress_count < p.num_q_tokens); - p.plan_c[compress_count++] = plan; - } else { - // trailing partial segment - RuntimeCheck(write_count < p.num_q_tokens); - p.plan_w[write_count++] = plan; - } - pos = seg_end; - } - counter += extend_len; - } - RuntimeCheck(counter == p.num_q_tokens, "input size ", counter, " != num_q_tokens ", p.num_q_tokens); - return std::tuple{compress_count, write_count}; -} - -struct OnlinePrefillStage1Params { - CompressPlan* __restrict__ plan_c; - CompressPlan* __restrict__ plan_w; - const int64_t* __restrict__ req_pool_indices; // (batch_size,) - const int32_t* __restrict__ req_to_token; // (num_reqs, max_tokens) - const int64_t* __restrict__ full_to_swa; // (full_cache_size,) - int64_t stride_r2t; - int32_t swa_page_size; - uint32_t num_c; - uint32_t num_w; -}; - -__global__ void plan_c128_online_prefill_kernel(const OnlinePrefillStage1Params params) { - const uint32_t idx = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t total = params.num_c + params.num_w; - if (idx >= total) return; - - const bool is_compress = idx < params.num_c; - CompressPlan* const plan_ptr = is_compress ? ¶ms.plan_c[idx] : ¶ms.plan_w[idx - params.num_c]; - auto plan = *plan_ptr; - const auto batch_id = plan.read_page_0; - const auto rid = params.req_pool_indices[batch_id]; - const int32_t position = static_cast(plan.seq_len - 1u); - const int32_t chunk_start = (position / 128) * 128; - const int32_t full_loc = params.req_to_token[rid * params.stride_r2t + chunk_start]; - const int32_t swa_loc = static_cast(params.full_to_swa[full_loc]); - plan.read_page_0 = swa_loc / params.swa_page_size; - *plan_ptr = plan; -} - -using OnlinePrefillPlan = tvm::ffi::Tuple; - -inline OnlinePrefillPlan plan_online_prefill( - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView extend_lens, - const tvm::ffi::TensorView req_pool_indices, - const tvm::ffi::TensorView req_to_token, - const tvm::ffi::TensorView full_to_swa, - const tvm::ffi::TensorView plan_c_pin, - const tvm::ffi::TensorView plan_w_pin, - const tvm::ffi::TensorView plan_c_dev_, - const tvm::ffi::TensorView plan_w_dev_, - const int32_t swa_page_size) { - auto B = SymbolicSize{"batch_size"}; - auto N = SymbolicSize{"num_q_tokens"}; - auto cpu = SymbolicDevice{}; - auto device_ = SymbolicDevice{}; - cpu.set_options(); - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(cpu) - .verify(seq_lens) - .verify(extend_lens); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({-1, -1}) // - .with_dtype() - .with_device(device_) - .verify(req_to_token); - TensorMatcher({-1}) // - .with_dtype() - .with_device(device_) - .verify(full_to_swa); - TensorMatcher({N, sizeof(CompressPlan)}) // - .with_dtype() - .with_device(cpu) - .verify(plan_c_pin) - .verify(plan_w_pin); - TensorMatcher({N, sizeof(CompressPlan)}) // - .with_dtype() - .with_device(device_) - .verify(plan_c_dev_) - .verify(plan_w_dev_); - - const auto stage0_params = OnlinePrefillStage0Params{ - .plan_c = static_cast(plan_c_pin.data_ptr()), - .plan_w = static_cast(plan_w_pin.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .extend_lens = static_cast(extend_lens.data_ptr()), - .batch_size = static_cast(B.unwrap()), - .num_q_tokens = static_cast(N.unwrap()), - }; - - // Debug instrumentation: SGLANG_DEBUG_C128_ONLINE_GUARD=1 wraps stage 0 - // with redzone + post-write magic-check on the pin buffers, plus a strict - // upper-bound check on `batch_size` and `num_q_tokens`. If stage 0 has a - // CPU OOB this trips a clear panic at the offending byte instead of a - // delayed CUDA IMA from corrupted heap memory. - static const bool kGuard = []() { - const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_GUARD"); - return v != nullptr && v[0] == '1'; - }(); - if (kGuard) { - RuntimeCheck(stage0_params.batch_size <= 65536u, "batch_size out of bound: ", stage0_params.batch_size); - RuntimeCheck(stage0_params.num_q_tokens <= 65536u, "num_q_tokens out of bound: ", stage0_params.num_q_tokens); - // Stamp the pin buffers with 0xAB so we can detect any byte still 0xAB - // beyond what stage 0 should have written (= OOB never reached, that's fine) - // or any byte BEYOND num_q_tokens*16 written to (= true OOB into - // adjacent allocation). - auto* pc = static_cast(plan_c_pin.data_ptr()); - auto* pw = static_cast(plan_w_pin.data_ptr()); - const auto bytes = static_cast(N.unwrap()) * sizeof(CompressPlan); - std::memset(pc, 0xAB, bytes); - std::memset(pw, 0xAB, bytes); - } - - const auto [num_c, num_w] = _plan_prefill_partial(stage0_params); - - if (kGuard) { - // Verify stage 0 wrote ONLY to the [0, num_c*16) and [0, num_w*16) prefix. - auto* pc = static_cast(plan_c_pin.data_ptr()); - auto* pw = static_cast(plan_w_pin.data_ptr()); - const auto end_c = static_cast(num_c) * sizeof(CompressPlan); - const auto end_w = static_cast(num_w) * sizeof(CompressPlan); - const auto pin_bytes = static_cast(N.unwrap()) * sizeof(CompressPlan); - for (size_t k = end_c; k < pin_bytes; ++k) { - RuntimeCheck( - pc[k] == 0xAB, - "GUARD: plan_c_pin OOB write at byte ", - k, - " (num_c=", - num_c, - ", num_q_tokens=", - N.unwrap(), - ")"); - } - for (size_t k = end_w; k < pin_bytes; ++k) { - RuntimeCheck( - pw[k] == 0xAB, - "GUARD: plan_w_pin OOB write at byte ", - k, - " (num_w=", - num_w, - ", num_q_tokens=", - N.unwrap(), - ")"); - } - } - - const auto device = device_.unwrap(); - // Out-params pre-allocated by Python. Cast to typed pointers for use. - auto* const plan_c_dev_ptr = static_cast(plan_c_dev_.data_ptr()); - auto* const plan_w_dev_ptr = static_cast(plan_w_dev_.data_ptr()); - - if (const auto total = num_c + num_w) { - const auto stream = LaunchKernel::resolve_device(device); - // SGLANG_DEBUG_C128_ONLINE_SYNC_H2D=1 forces a synchronous H2D copy. - static const bool kSyncH2D = []() { - const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_SYNC_H2D"); - return v != nullptr && v[0] == '1'; - }(); - // SGLANG_DEBUG_C128_ONLINE_NO_H2D=1 skips the H2D copy entirely (debug only). - static const bool kNoH2D = []() { - const char* v = std::getenv("SGLANG_DEBUG_C128_ONLINE_NO_H2D"); - return v != nullptr && v[0] == '1'; - }(); - const auto copy_to_device = [stream](void* dst, void* src, int64_t count) { - if (kNoH2D) return; - const auto bytes = count * sizeof(CompressPlan); - if (kSyncH2D) { - RuntimeDeviceCheck(::cudaMemcpy(dst, src, bytes, ::cudaMemcpyHostToDevice)); - } else { - RuntimeDeviceCheck(::cudaMemcpyAsync(dst, src, bytes, ::cudaMemcpyHostToDevice, stream)); - } - }; - if (num_c) copy_to_device(plan_c_dev_ptr, plan_c_pin.data_ptr(), num_c); - if (num_w) copy_to_device(plan_w_dev_ptr, plan_w_pin.data_ptr(), num_w); - - const auto stage1_params = OnlinePrefillStage1Params{ - .plan_c = plan_c_dev_ptr, - .plan_w = plan_w_dev_ptr, - .req_pool_indices = static_cast(req_pool_indices.data_ptr()), - .req_to_token = static_cast(req_to_token.data_ptr()), - .full_to_swa = static_cast(full_to_swa.data_ptr()), - .stride_r2t = req_to_token.stride(0), - .swa_page_size = swa_page_size, - .num_c = num_c, - .num_w = num_w, - }; - constexpr uint32_t kBlockSize = 128; - const auto num_blocks = host::div_ceil(total, kBlockSize); - LaunchKernel(num_blocks, kBlockSize, device)(plan_c128_online_prefill_kernel, stage1_params); - } - return OnlinePrefillPlan{num_c, num_w}; -} - -} // namespace host::compress - -namespace { - -[[maybe_unused]] -constexpr auto& plan_compress_128_online_decode = host::compress::plan_online_decode; -[[maybe_unused]] -constexpr auto& plan_compress_128_online_prefill = host::compress::plan_online_prefill; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh deleted file mode 100644 index 31353e6a15..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c128_v2.cuh +++ /dev/null @@ -1,448 +0,0 @@ -/** - * \brief Here's some dimension info for the main buffer used in C128 prefill and decode. - * - * kv_buffer: [num_indices, 128, head_dim * 2] - * - last dimension layout: | kv | score | - * kv_input: [batch_size, head_dim * 2] - * kv_output: [batch_size, head_dim] - * score_bias (ape): [128, head_dim] - * plan_c/plan_w: [variable length] - * - * For prefill, batch_size = num_q_tokens - */ - -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include - -namespace { - -using PlanD = device::compress::DecodePlan; -using PlanC = device::compress::CompressPlan; -using PlanW = device::compress::WritePlan; - -/// \brief Each thread will handle this many elements (split along head_dim) -constexpr int32_t kTileElements = 2; -/// \brief Each warp will handle this many elements (split along 128) -constexpr int32_t kElementsPerWarp = 8; -constexpr uint32_t kNumWarps = 128 / kElementsPerWarp; -constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps; -constexpr uint32_t kWriteBlockSize = 128; // one warp per write - -/// \brief Need to reduce register usage to increase occupancy -#define C128_KERNEL __global__ __launch_bounds__(kBlockSize, 2) -#define WRITE_KERNEL __global__ __launch_bounds__(kWriteBlockSize, 16) - -struct Compress128DecodeParams { - void* __restrict__ kv_buffer; - const void* __restrict__ kv_input; - void* __restrict__ kv_output; - const void* __restrict__ score_bias; - const PlanD* __restrict__ plan_d; - uint32_t batch_size; -}; - -struct Compress128PrefillParams { - void* __restrict__ kv_buffer; - const void* __restrict__ kv_input; - void* __restrict__ kv_output; - const void* __restrict__ score_bias; - const PlanC* __restrict__ plan_c; - const PlanW* __restrict__ plan_w; - uint32_t num_compress; - uint32_t num_write; -}; - -struct Compress128SharedBuffer { - using Storage = device::AlignedVector; - Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict - SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) { - return data[warp_id][lane_id]; - } - SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) { - return data[warp_id][lane_id][tile_id]; - } -}; - -template -struct C128Trait { - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 - static constexpr int64_t kHeadDim = kHeadDim_; - static constexpr int64_t kScoreOffset = kHeadDim; - static constexpr int64_t kElementSize = kHeadDim * 2; - static constexpr int64_t kPageElementSize = 128 * kElementSize; // page size = 128 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static_assert(kHeadDim % kTileDim == 0); -}; - -template -SGL_DEVICE void c128_forward( - const InFloat* kv_buf, // [128n, 128n + 127] - const InFloat* kv_src, // ragged pointer at position = 128n + 127 - OutFloat* kv_out, - const InFloat* score_bias, - const int32_t buffer_len) { - using namespace device; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - - /// NOTE: part 1: load kv + score - using StorageIn = AlignedVector; - const auto gmem_in = tile::Memory{lane_id, kWarpThreads}; - StorageIn kv[kElementsPerWarp]; - StorageIn score[kElementsPerWarp]; - StorageIn bias[kElementsPerWarp]; - const int32_t warp_offset = warp_id * kElementsPerWarp; - -#pragma unroll - for (int32_t i = 0; i < 8; ++i) { - const int32_t j = i + warp_offset; - bias[i] = gmem_in.load(score_bias + j * Trait::kHeadDim); - } - - const auto kv_start = kv_src - 127 * Trait::kElementSize; // point to start - -#pragma unroll - for (int32_t i = 0; i < kElementsPerWarp; ++i) { - const int32_t j = i + warp_offset; - __builtin_assume(j < 128); - const auto src = j < buffer_len ? kv_buf : kv_start; - kv[i] = gmem_in.load(src + j * Trait::kElementSize); - score[i] = gmem_in.load(src + j * Trait::kElementSize + Trait::kScoreOffset); - } - - /// NOTE: part 2: safe online softmax + weighted sum - using TmpStorage = typename Compress128SharedBuffer::Storage; - __shared__ Compress128SharedBuffer s_local_val_max; - __shared__ Compress128SharedBuffer s_local_exp_sum; - __shared__ Compress128SharedBuffer s_local_product; - - TmpStorage tmp_val_max; - TmpStorage tmp_exp_sum; - TmpStorage tmp_product; - - float score_fp32[kTileElements][kElementsPerWarp]; - - // convert to fp32 and apply bias first -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - for (int32_t j = 0; j < kElementsPerWarp; ++j) { - score_fp32[i][j] = cast(score[j][i]) + cast(bias[j][i]); - } - } - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - const auto& score = score_fp32[i]; - float max_value = score[0]; - float sum_exp_value = 0.0f; - -#pragma unroll - for (int32_t j = 1; j < kElementsPerWarp; ++j) { - const auto fp32_score = score[j]; - max_value = fmaxf(max_value, fp32_score); - } - - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - const auto fp32_score = score[j]; - const auto exp_score = expf(fp32_score - max_value); - sum_product += cast(kv[j][i]) * exp_score; - sum_exp_value += exp_score; - } - - tmp_val_max[i] = max_value; - tmp_exp_sum[i] = sum_exp_value; - tmp_product[i] = sum_product; - } - - // naturally aligned, so no bank conflict - s_local_val_max(warp_id, lane_id) = tmp_val_max; - s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum; - s_local_product(warp_id, lane_id) = tmp_product; - - __syncthreads(); - - /// NOTE: part 3: online softmax - /// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce - /// each reduce will consume `kNumWarps` threads (use partial warp reduction) - constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps; - constexpr uint32_t kIteration = kReductionCount / kBlockSize; - - PDLTriggerSecondary(); - -#pragma unroll - for (uint32_t i = 0; i < kIteration; ++i) { - /// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)` - const uint32_t j = i * kBlockSize + warp_id * kWarpThreads + lane_id; - /// NOTE: Range `[0, kNumWarps)` - const uint32_t local_warp_id = j % kNumWarps; - /// NOTE: Range `[0, kTileElements * kWarpThreads)` - const uint32_t local_elem_id = j / kNumWarps; - /// NOTE: Range `[0, kTileElements)` - const uint32_t local_tile_id = local_elem_id % kTileElements; - /// NOTE: Range `[0, kWarpThreads)` - const uint32_t local_lane_id = local_elem_id / kTileElements; - /// NOTE: each warp will access the whole tile (all `kTileElements`) - /// and for different lanes, the memory access only differ in `local_warp_id` - /// so there's no bank conflict in shared memory access. - static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs"); - const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id); - const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id); - const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id); - const auto global_val_max = warp::reduce_max(local_val_max); - const auto rescale = expf(local_val_max - global_val_max); - const auto global_exp_sum = warp::reduce_sum(local_exp_sum * rescale); - const auto final_scale = rescale / global_exp_sum; - const auto global_product = warp::reduce_sum(local_product * final_scale); - kv_out[local_elem_id] = cast(global_product); - } -} - -template -SGL_DEVICE void c128_write_decode(InFloat* kv_buf, const InFloat* kv_src) { - using namespace device; - - using Storage = AlignedVector; - const auto gmem = tile::Memory::warp(); - - Storage data[2]; -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - data[i] = gmem.load(kv_src + Trait::kHeadDim * i); - } -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - gmem.store(kv_buf + Trait::kHeadDim * i, data[i]); - } -} - -template -C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) { - using namespace device; - using Trait = C128Trait; - - const uint32_t warp_id = threadIdx.x / kWarpThreads; - const uint32_t global_bid = blockIdx.x / Trait::kNumSplit; // batch id - const uint32_t global_sid = blockIdx.x % Trait::kNumSplit; // split id - const int64_t split_offset = global_sid * Trait::kTileDim; - if (global_bid >= params.batch_size) return; - - const auto plan = params.plan_d[global_bid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_output = static_cast(params.kv_output) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - const auto score_bias = static_cast(params.score_bias) + split_offset; - - const auto kv_src = kv_input + global_bid * Trait::kElementSize; - const auto kv_out = kv_output + global_bid * Trait::kHeadDim; - const auto kv_buf = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; - const auto kv_dst = kv_buffer + plan.write_loc * Trait::kElementSize; - - PDLWaitPrimary(); - // the write warp must match the load warp in the following `c128_forward` - if (warp_id == kNumWarps - 1) { - c128_write_decode(kv_dst, kv_src); - } - if (plan.write_loc % 128 == 127) { - c128_forward(kv_buf, kv_src, kv_out, score_bias, 128); - } -} - -// compress kernel -template -C128_KERNEL void flash_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { - using namespace device; - using Trait = C128Trait; - - const uint32_t global_pid = blockIdx.x / Trait::kNumSplit; // plan id - const uint32_t global_sid = blockIdx.x % Trait::kNumSplit; // split id - const int64_t split_offset = global_sid * Trait::kTileDim; - if (global_pid >= params.num_compress) return; - - const auto plan = params.plan_c[global_pid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_output = static_cast(params.kv_output) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - const auto score_bias = static_cast(params.score_bias) + split_offset; - if (plan.is_invalid()) return; - - const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; - // Compact output: one row per compress plan, indexed by `global_pid`. - const auto kv_out = kv_output + global_pid * Trait::kHeadDim; - const auto kv_buf = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; - PDLWaitPrimary(); - c128_forward(kv_buf, kv_src, kv_out, score_bias, plan.buffer_len); -} - -template -WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillParams params) { - using namespace device; - using Trait = C128Trait; - using StorageIn = AlignedVector; - - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id - const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id - // split the contiguous `kHeadDim * 2` into `kNumSplit` tiles - // each warp handles 1 contiguous tile (in contrast, decode handle the strided head_dim) - const int64_t split_offset = global_sid * (Trait::kTileDim * 2); - if (global_pid >= params.num_write) return; - - const auto plan = params.plan_w[global_pid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - if (plan.is_invalid()) return; - - // each warp will handle a contiguous region - const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; - const auto kv_buf = kv_buffer + plan.write_loc * Trait::kElementSize; - const auto gmem = tile::Memory::warp(); - - PDLWaitPrimary(); - StorageIn data[2]; -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - data[i] = gmem.load(kv_src, i); - } - PDLTriggerSecondary(); -#pragma unroll - for (int32_t i = 0; i < 2; ++i) { - gmem.store(kv_buf, data[i], i); - } -} - -template -struct FlashCompress128Kernel { - static constexpr auto decode_kernel = flash_c128_decode; - static constexpr auto prefill_c_kernel = flash_c128_prefill; - static constexpr auto prefill_w_kernel = write_c128_prefill; - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - using Trait = C128Trait; - - static void run_decode( - const tvm::ffi::TensorView kv_buffer, - const tvm::ffi::TensorView kv_input, - const tvm::ffi::TensorView kv_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_d_) { - using namespace host; - - auto N = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 128, Trait::kElementSize}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_buffer); - TensorMatcher({N, Trait::kElementSize}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_input); - TensorMatcher({N, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - - const auto plan_d = compress::verify_plan_d(plan_d_, N, device_); - const auto batch_size = static_cast(N.unwrap()); - const auto params = Compress128DecodeParams{ - .kv_buffer = kv_buffer.data_ptr(), - .kv_input = kv_input.data_ptr(), - .kv_output = kv_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_d = plan_d, - .batch_size = batch_size, - }; - const uint32_t num_blocks = batch_size * kNumSplit; - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(decode_kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_buffer, - const tvm::ffi::TensorView kv_input, - const tvm::ffi::TensorView kv_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_c_, - const tvm::ffi::TensorView plan_w_) { - using namespace host; - - auto N = SymbolicSize{"num_q_tokens"}; - auto C = SymbolicSize{"num_c_plans"}; - auto W = SymbolicSize{"num_w_plans"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 128, Trait::kElementSize}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_buffer); - TensorMatcher({N, Trait::kElementSize}) // kv score input (ragged) - .with_dtype() - .with_device(device_) - .verify(kv_input); - TensorMatcher({C, kHeadDim}) // kv compressed output (compact) - .with_dtype() - .with_device(device_) - .verify(kv_output); - TensorMatcher({128, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - - const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); - const auto plan_w = compress::verify_plan_w(plan_w_, W, device_); - const auto device = device_.unwrap(); - const auto num_q_tokens = static_cast(N.unwrap()); - const auto num_c = static_cast(C.unwrap()); - const auto num_w = static_cast(W.unwrap()); - const auto params = Compress128PrefillParams{ - .kv_buffer = kv_buffer.data_ptr(), - .kv_input = kv_input.data_ptr(), - .kv_output = kv_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_c = plan_c, - .plan_w = plan_w, - .num_compress = num_c, - .num_write = num_w, - }; - RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); - if (const auto num_c_blocks = num_c * kNumSplit) { - constexpr auto kBlockSize_C = kBlockSize; - LaunchKernel(num_c_blocks, kBlockSize_C, device) // - .enable_pdl(kUsePDL)(prefill_c_kernel, params); - } - constexpr uint32_t kWarpsPerWriteBlock = kWriteBlockSize / device::kWarpThreads; - if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerWriteBlock)) { - constexpr auto kBlockSize_W = kWriteBlockSize; - LaunchKernel(num_w_blocks, kBlockSize_W, device) // - .enable_pdl(kUsePDL)(prefill_w_kernel, params); - } - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh deleted file mode 100644 index 145ab1fb08..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4.cuh +++ /dev/null @@ -1,549 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include - -namespace { - -using Plan4 = device::compress::PrefillPlan; -using IndiceT = int32_t; - -/// \brief Each thread will handle this many elements (split along head_dim) -constexpr int kTileElements = 4; - -/// \brief Need to improve register usage to reduce latency -#define C4_KERNEL __global__ __launch_bounds__(128, 4) - -enum class PageMode { - RingBuffer = 8, - Page4Align = 4, -}; - -struct alignas(16) C4IndexBundle { - int32_t load_first_page; - int32_t load_second_page; - int32_t write_first_page; - int32_t last_position; -}; - -struct Compress4DecodeParams { - /** - * \brief Shape: `[num_indices, 8, head_dim * 4]` \n - * last dimension layout: - * | kv overlap | kv | score overlap | score | - */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[batch_size, head_dim * 4]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[batch_size, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[8, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, ]` */ - const IndiceT* __restrict__ seq_lens; - /** \brief Shape: `[batch_size, 1]` */ - const int32_t* __restrict__ extra; - /** \NOTE: `batch_size` <= `num_indices` */ - uint32_t batch_size; -}; - -struct Compress4PrefillParams { - /** - * \brief Shape: `[num_indices, 8, head_dim * 4]` \n - * last dimension layout: - * | kv overlap | kv | score overlap | score | - */ - void* __restrict__ kv_score_buffer; - /** \brief Shape: `[num_q_tokens, head_dim * 4]` */ - const void* __restrict__ kv_score_input; - /** \brief Shape: `[num_q_tokens, head_dim]` */ - void* __restrict__ kv_compressed_output; - /** \brief Shape: `[8, head_dim]` (called `ape`) */ - const void* __restrict__ score_bias; - /** \brief Shape: `[batch_size, ]`*/ - const IndiceT* __restrict__ indices; - /** \brief Shape: `[batch_size, 4]` */ - const C4IndexBundle* __restrict__ extra; - /** \brief The following part is plan info. */ - - const Plan4* __restrict__ compress_plan; - const Plan4* __restrict__ write_plan; - uint32_t num_compress; - uint32_t num_write; -}; - -template -SGL_DEVICE void c4_write( - T* kv_score_buf, // - const T* kv_score_src, - const int64_t head_dim, - const int32_t write_pos) { - using namespace device; - - using Storage = AlignedVector; - const auto element_size = head_dim * 4; - const auto gmem = tile::Memory::warp(); - kv_score_buf += write_pos * element_size; - - /// NOTE: Layout | [0] = kv overlap | [1] = kv | [2] = score overlap | [3] = score | - Storage kv_score[4]; -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - kv_score[i] = gmem.load(kv_score_src + head_dim * i); - } -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - gmem.store(kv_score_buf + head_dim * i, kv_score[i]); - } -} - -template -SGL_DEVICE void c4_forward( - const InFloat* kv_score_buf, - const InFloat* kv_score_src, - OutFloat* kv_out, - const InFloat* score_bias, - const int64_t head_dim, - const int32_t seq_len, - const int32_t window_len, - [[maybe_unused]] const InFloat* kv_score_overlap_buf = nullptr) { - using namespace device; - - const auto element_size = head_dim * 4; - const auto score_offset = head_dim * 2; - const auto overlap_stride = head_dim; - - /// NOTE: part 1: load kv + score - using StorageIn = AlignedVector; - const auto gmem_in = tile::Memory::warp(); - StorageIn kv[8]; - StorageIn score[8]; - StorageIn bias[8]; - -#pragma unroll - for (int32_t i = 0; i < 8; ++i) { - bias[i] = gmem_in.load(score_bias + i * head_dim); - } - -#pragma unroll - for (int32_t i = 0; i < 8; ++i) { - const bool is_overlap = i < 4; - const InFloat* src; - if (i < window_len) { - /// NOTE: `seq_len` must be a multiple of 4 here - if constexpr (kPaged) { - const auto kv_score_ptr = is_overlap ? kv_score_overlap_buf : kv_score_buf; - const int32_t k = i % 4; - src = kv_score_ptr + k * element_size; - } else { - const int32_t k = (seq_len + i) % 8; - src = kv_score_buf + k * element_size; - } - } else { - /// NOTE: k in [-7, 0]. We'll load from the ragged `kv_score_src` - const int32_t k = i - 7; - src = kv_score_src + k * element_size; - } - src += (is_overlap ? 0 : overlap_stride); - kv[i] = gmem_in.load(src); - score[i] = gmem_in.load(src + score_offset); - } - - if (seq_len == 4) { - [[unlikely]]; - constexpr float kFloatNegInf = -1e9f; -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - kv[i].fill(cast(0.0f)); - score[i].fill(cast(kFloatNegInf)); - } - } - - /// NOTE: part 2: safe online softmax + weighted sum - using StorageOut = AlignedVector; - const auto gmem_out = tile::Memory::warp(); - StorageOut result; - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - float score_fp32[8]; - -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - score_fp32[j] = cast(score[j][i]) + cast(bias[j][i]); - } - - float max_value = score_fp32[0]; - float sum_exp_value = 0.0f; - -#pragma unroll - for (int32_t j = 1; j < 8; ++j) { - const auto fp32_score = score_fp32[j]; - max_value = fmaxf(max_value, fp32_score); - } - - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - const auto fp32_score = score_fp32[j]; - const auto exp_score = expf(fp32_score - max_value); - sum_product += cast(kv[j][i]) * exp_score; - sum_exp_value += exp_score; - } - - result[i] = cast(sum_product / sum_exp_value); - } - - gmem_out.store(kv_out, result); -} - -template -C4_KERNEL void flash_c4_decode(const __grid_constant__ Compress4DecodeParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 128 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - constexpr int64_t kElementSize = kHeadDim * 4; // `* 4` due to overlap transform + score - static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); - - const auto& [ - _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score - indices, seq_lens, extra, batch_size // decode info - ] = params; - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_bid = global_wid / kNumSplit; // batch id - const uint32_t global_sid = global_wid % kNumSplit; // split id - - if (global_bid >= batch_size) return; - - const int32_t index = indices[global_bid]; - const int32_t seq_len = seq_lens[global_bid]; - const int64_t split_offset = global_sid * kTileDim; - - // kv score - const auto kv_score_buffer = static_cast(_kv_score_buffer); - - // kv input - const auto kv_score_input = static_cast(_kv_score_input); - const auto kv_src = kv_score_input + global_bid * kElementSize + split_offset; - - // kv output - const auto kv_compressed_output = static_cast(_kv_compressed_output); - const auto kv_out = kv_compressed_output + global_bid * kHeadDim + split_offset; - - // score bias (ape) - const auto score_bias = static_cast(_score_bias) + split_offset; - - PDLWaitPrimary(); - - /// NOTE: `position` = `seq_len - 1`. To avoid underflow, we use `seq_len + page_size - 1` - if constexpr (kMode == PageMode::Page4Align) { - const auto index_prev = extra[global_bid]; - const auto kv_buf = kv_score_buffer + index * (kElementSize * 4) + split_offset; - c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 3) % 4); - if (seq_len % 4 == 0) { - const auto kv_overlap = kv_buf + (index_prev - index) * (kElementSize * 4); - c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, 8, kv_overlap); - } - } else { - static_assert(kMode == PageMode::RingBuffer, "Unsupported PageMode"); - const auto kv_buf = kv_score_buffer + index * (kElementSize * 8) + split_offset; - c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/(seq_len + 7) % 8); - if (seq_len % 4 == 0) { - c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, /*window_size=*/8); - } - } - - PDLTriggerSecondary(); -} - -template -C4_KERNEL void flash_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { - using namespace device; - - constexpr int64_t kTileDim = kTileElements * kWarpThreads; // 128 - constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - constexpr int64_t kElementSize = kHeadDim * 4; // `* 4` due to overlap transform + score - static_assert(kHeadDim % kTileDim == 0, "Head dim must be multiple of tile dim"); - - const auto& [ - _kv_score_buffer, _kv_score_input, _kv_compressed_output, _score_bias, // kv score - indices, extra, compress_plan, write_plan, num_compress, num_write // prefill plan - ] = params; - - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_pid = global_wid / kNumSplit; // plan id - const uint32_t global_sid = global_wid % kNumSplit; // split id - - /// NOTE: compiler can optimize this if-else at compile time - const auto num_plans = kWrite ? num_write : num_compress; - const auto plan_ptr = kWrite ? write_plan : compress_plan; - if (global_pid >= num_plans) return; - - const auto& [ragged_id, global_bid, position, window_len] = plan_ptr[global_pid]; - const int64_t split_offset = global_sid * kTileDim; - - // kv score - const auto kv_score_buffer = static_cast(_kv_score_buffer); - - // kv input - const auto kv_score_input = static_cast(_kv_score_input); - const auto kv_src = kv_score_input + ragged_id * kElementSize + split_offset; - - // kv output - const auto kv_compressed_output = static_cast(_kv_compressed_output); - const auto kv_out = kv_compressed_output + ragged_id * kHeadDim + split_offset; - - if (ragged_id == 0xFFFFFFFF) [[unlikely]] - return; - - // score bias (ape) - const auto score_bias = static_cast(_score_bias) + split_offset; - const auto seq_len = position + 1; - const int32_t index = indices[global_bid]; - - PDLWaitPrimary(); - - if constexpr (kMode == PageMode::Page4Align) { - const auto write_second_page = index; - const auto [load_first_page, load_second_page, write_first_page, last_pos] = extra[global_bid]; - if constexpr (kWrite) { - int32_t index; - if (position < static_cast(last_pos)) { - index = write_first_page; - } else { - index = write_second_page; - } - const auto kv_buf = kv_score_buffer + index * (kElementSize * 4) + split_offset; - c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 4); - } else { - int32_t index_overlap, index_normal; - if (window_len <= 4) { - index_overlap = load_second_page; - index_normal = load_second_page; // not used - } else { - index_overlap = load_first_page; - index_normal = load_second_page; - } - const auto kv_buf = kv_score_buffer + index_normal * (kElementSize * 4) + split_offset; - const auto kv_overlap = kv_score_buffer + index_overlap * (kElementSize * 4) + split_offset; - c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, window_len, kv_overlap); - } - } else { - static_assert(kMode == PageMode::RingBuffer, "Unsupported PageMode"); - const auto kv_buf = kv_score_buffer + index * (kElementSize * 8) + split_offset; - if constexpr (kWrite) { - c4_write(kv_buf, kv_src, kHeadDim, /*write_pos=*/position % 8); - } else { - c4_forward(kv_buf, kv_src, kv_out, score_bias, kHeadDim, seq_len, window_len); - } - } - - PDLTriggerSecondary(); -} - -template -struct FlashCompress4Kernel { - template - static constexpr auto decode_kernel = flash_c4_decode; - template - static constexpr auto prefill_kernel = flash_c4_prefill; - template - static constexpr auto prefill_c_kernel = prefill_kernel; - template - static constexpr auto prefill_w_kernel = prefill_kernel; - static constexpr uint32_t kBlockSize = 128; - static constexpr uint32_t kTileDim = kTileElements * device::kWarpThreads; - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static constexpr uint32_t kWarpsPerBlock = kBlockSize / device::kWarpThreads; - - using Self = FlashCompress4Kernel; - - static void run_decode( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::Optional extra) { - using namespace host; - - // this should not happen in practice - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - const auto extra_ptr = _get_extra_pointer(B, device_, extra); - const auto page_size = extra_ptr != nullptr ? 4 : 8; - - TensorMatcher({-1, page_size, kHeadDim * 4}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({B, kHeadDim * 4}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({B, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({8, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - TensorMatcher({B}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - TensorMatcher({B}) // seq lens - .with_dtype() - .with_device(device_) - .verify(seq_lens); - - const auto device = device_.unwrap(); - const auto batch_size = static_cast(B.unwrap()); - const auto params = Compress4DecodeParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .extra = static_cast(extra_ptr), - .batch_size = batch_size, - }; - const auto kernel = extra_ptr != nullptr ? decode_kernel // - : decode_kernel; - const uint32_t num_blocks = div_ceil(batch_size * kNumSplit, kWarpsPerBlock); - LaunchKernel(num_blocks, kBlockSize, device) // - .enable_pdl(kUsePDL)(kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_score_buffer, - const tvm::ffi::TensorView kv_score_input, - const tvm::ffi::TensorView kv_compressed_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView indices, - const tvm::ffi::TensorView compress_plan, - const tvm::ffi::TensorView write_plan, - const tvm::ffi::Optional extra) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto N = SymbolicSize{"num_q_tokens"}; - auto X = SymbolicSize{"compress_tokens"}; - auto Y = SymbolicSize{"write_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - const auto extra_ptr = _get_extra_pointer(B, device_, extra, /*is_prefill=*/true); - const auto page_size = extra_ptr != nullptr ? 4 : 8; - - TensorMatcher({-1, page_size, kHeadDim * 4}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_score_buffer); - TensorMatcher({N, kHeadDim * 4}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_score_input); - TensorMatcher({N, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_compressed_output); - TensorMatcher({8, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - TensorMatcher({B}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - TensorMatcher({X, compress::kPrefillPlanDim}) // compress plan - .with_dtype() - .with_device(device_) - .verify(compress_plan); - TensorMatcher({Y, compress::kPrefillPlanDim}) // write plan - .with_dtype() - .with_device(device_) - .verify(write_plan); - - const auto device = device_.unwrap(); - const auto batch_size = static_cast(B.unwrap()); - const auto num_q_tokens = static_cast(N.unwrap()); - const auto num_c = static_cast(X.unwrap()); - const auto num_w = static_cast(Y.unwrap()); - const auto params = Compress4PrefillParams{ - .kv_score_buffer = kv_score_buffer.data_ptr(), - .kv_score_input = kv_score_input.data_ptr(), - .kv_compressed_output = kv_compressed_output.data_ptr(), - .score_bias = ape.data_ptr(), - .indices = static_cast(indices.data_ptr()), - .extra = static_cast(extra_ptr), - .compress_plan = static_cast(compress_plan.data_ptr()), - .write_plan = static_cast(write_plan.data_ptr()), - .num_compress = num_c, - .num_write = num_w, - }; - RuntimeCheck(num_q_tokens >= batch_size, "num_q_tokens must be >= batch_size"); - RuntimeCheck(num_q_tokens >= std::max(num_c, num_w), "invalid prefill plan"); - if (const auto num_c_blocks = div_ceil(num_c * kNumSplit, kWarpsPerBlock)) { - const auto c_kernel = extra_ptr != nullptr ? prefill_c_kernel // - : prefill_c_kernel; - LaunchKernel(num_c_blocks, kBlockSize, device) // - .enable_pdl(kUsePDL)(c_kernel, params); - } - if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerBlock)) { - const auto w_kernel = extra_ptr != nullptr ? prefill_w_kernel // - : prefill_w_kernel; - LaunchKernel(num_w_blocks, kBlockSize, device) // - .enable_pdl(kUsePDL)(w_kernel, params); - } - } - - // some auxiliary functions - private: - static const void* _get_extra_pointer( - host::SymbolicSize& B, // batch_size - host::SymbolicDevice& device, - const tvm::ffi::Optional& extra, - bool is_prefill = false) { - // only have value when using page-aligned mode - if (!extra.has_value()) return nullptr; - const auto& extra_tensor = extra.value(); - /// NOTE: the metadata layout is different for prefill and decode: - /// for prefill, last 4 are: - /// load overlap | load normal | write overlap | last written page - /// for decode, last 1 is the write (also load) overlap - host::TensorMatcher({B, is_prefill ? 4 : 1}) // extra tensor - .with_dtype() - .with_device(device) - .verify(extra_tensor); - const auto data_ptr = extra_tensor.data_ptr(); - host::RuntimeCheck(data_ptr != nullptr, "extra tensor data ptr is null"); - if (is_prefill) { - static_assert(alignof(C4IndexBundle) == 16); - host::RuntimeCheck(std::bit_cast(data_ptr) % 16 == 0, "extra tensor is not properly aligned"); - } - return data_ptr; - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh deleted file mode 100644 index efa9f05100..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c4_v2.cuh +++ /dev/null @@ -1,405 +0,0 @@ -/** - * \brief Here's some dimension info for the main buffer used in C4 prefill and decode. - * - * kv_buffer: [num_indices, 8, head_dim * 4] - * - last dimension layout: | kv overlap | kv | score overlap | score | - * kv_input: [batch_size, head_dim * 4] - * kv_output: [batch_size, head_dim] - * score_bias (ape): [8, head_dim] - * plan_c/plan_w: [variable length] - * - * For prefill, batch_size = num_q_tokens - */ - -#include -#include - -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -#include -#include - -namespace { - -using PlanD = device::compress::DecodePlan; -using PlanC = device::compress::CompressPlan; -using PlanW = device::compress::WritePlan; - -/// \brief Each thread will handle this many elements (split along head_dim) -constexpr int32_t kTileElements = 4; - -/// \brief Need to improve register usage to reduce latency -#define C4_KERNEL __global__ __launch_bounds__(128, 4) -#define WRITE_KERNEL __global__ __launch_bounds__(128, 16) - -struct Compress4DecodeParams { - void* __restrict__ kv_buffer; - const void* __restrict__ kv_input; - void* __restrict__ kv_output; - const void* __restrict__ score_bias; - const PlanD* __restrict__ plan_d; - uint32_t batch_size; -}; - -struct Compress4PrefillParams { - void* __restrict__ kv_buffer; - const void* __restrict__ kv_input; - void* __restrict__ kv_output; - const void* __restrict__ score_bias; - const PlanC* __restrict__ plan_c; - const PlanW* __restrict__ plan_w; - uint32_t num_compress; - uint32_t num_write; -}; - -template -struct C4Trait { - static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 128 - static constexpr int64_t kHeadDim = kHeadDim_; - static constexpr int64_t kOverlapOffset = kHeadDim; - static constexpr int64_t kScoreOffset = kHeadDim * 2; - static constexpr int64_t kElementSize = kHeadDim * 4; - static constexpr int64_t kPageElementSize = 4 * kElementSize; // page size = 4 - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static_assert(kHeadDim % kTileDim == 0); -}; - -template -SGL_DEVICE void c4_forward( - const InFloat* kv_buf_0, // overlap [4n - 4, 4n - 1] - const InFloat* kv_buf_1, // normal [4n + 0, 4n + 3] - const InFloat* kv_src, // ragged pointer at position = 4n + 3 - OutFloat* kv_out, - const InFloat* score_bias, - const bool should_overlap, - const int32_t buffer_len) { - using namespace device; - - /// NOTE: part 1: load kv + score - using StorageIn = AlignedVector; - /// NOTE: load one tile_dim (< head_dim) at at time - const auto gmem_in = tile::Memory::warp(); - StorageIn kv[8]; - StorageIn score[8]; - StorageIn bias[8]; - -#pragma unroll - for (int32_t i = 0; i < 8; ++i) { - bias[i] = gmem_in.load(score_bias + i * Trait::kHeadDim); - } - - if (should_overlap) { - const auto kv_start = kv_src - 7 * Trait::kElementSize; // point to start -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - const auto src = i < buffer_len ? kv_buf_0 : kv_start; - const auto base = src + i * Trait::kElementSize; - kv[i] = gmem_in.load(base); - score[i] = gmem_in.load(base + Trait::kScoreOffset); - } - } else { - [[unlikely]]; - constexpr float kFloatNegInf = -FLT_MAX; -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - kv[i].fill(cast(0.0f)); - score[i].fill(cast(kFloatNegInf)); - } - } - - const auto kv_start = kv_src - 3 * Trait::kElementSize; // point to start -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - const auto src = i + 4 < buffer_len ? kv_buf_1 : kv_start; - const auto base = src + i * Trait::kElementSize + Trait::kOverlapOffset; - kv[i + 4] = gmem_in.load(base); - score[i + 4] = gmem_in.load(base + Trait::kScoreOffset); - } - - /// NOTE: part 2: safe online softmax + weighted sum - using StorageOut = AlignedVector; - const auto gmem_out = tile::Memory::warp(); - StorageOut result; - - // consume 32 fp registers - float score_fp32[kTileElements][8]; - - // convert to fp32 and apply bias first -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - for (int32_t j = 0; j < 8; ++j) { - score_fp32[i][j] = cast(score[j][i]) + cast(bias[j][i]); - } - } - -#pragma unroll - for (int32_t i = 0; i < kTileElements; ++i) { - const auto& score = score_fp32[i]; - float max_value = score[0]; - float sum_exp_value = 0.0f; - -#pragma unroll - for (int32_t j = 1; j < 8; ++j) { - const auto fp32_score = score[j]; - max_value = fmaxf(max_value, fp32_score); - } - - float sum_product = 0.0f; -#pragma unroll - for (int32_t j = 0; j < 8; ++j) { - const auto fp32_score = score[j]; - const auto exp_score = expf(fp32_score - max_value); - sum_product += cast(kv[j][i]) * exp_score; - sum_exp_value += exp_score; - } - - result[i] = cast(sum_product / sum_exp_value); - } - - // overlap the store with the next iteration's load - PDLTriggerSecondary(); - gmem_out.store(kv_out, result); -} - -template -SGL_DEVICE void c4_write_decode(InFloat* kv_buf, const InFloat* kv_src) { - using namespace device; - - using StorageIn = AlignedVector; - const auto gmem = tile::Memory::warp(); - - StorageIn data[4]; -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - data[i] = gmem.load(kv_src + Trait::kHeadDim * i); - } -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - gmem.store(kv_buf + Trait::kHeadDim * i, data[i]); - } -} - -template -C4_KERNEL void flash_c4_decode(const __grid_constant__ Compress4DecodeParams params) { - using namespace device; - using Trait = C4Trait; - - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_bid = global_wid / Trait::kNumSplit; // batch id - const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id - const int64_t split_offset = global_sid * Trait::kTileDim; - if (global_bid >= params.batch_size) return; - - const auto plan = params.plan_d[global_bid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_output = static_cast(params.kv_output) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - const auto score_bias = static_cast(params.score_bias) + split_offset; - - const auto kv_src = kv_input + global_bid * Trait::kElementSize; - const auto kv_out = kv_output + global_bid * Trait::kHeadDim; - const auto kv_buf_0 = kv_buffer + plan.read_page_0 * Trait::kPageElementSize; - const auto kv_buf_1 = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; - const auto kv_dst = kv_buffer + plan.write_loc * Trait::kElementSize; - - PDLWaitPrimary(); - c4_write_decode(kv_dst, kv_src); - if (plan.seq_len % 4 == 0) { - const auto need_overlap = plan.seq_len > 4; - c4_forward(kv_buf_0, kv_buf_1, kv_src, kv_out, score_bias, need_overlap, 8); - } -} - -template -C4_KERNEL void flash_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { - using namespace device; - using Trait = C4Trait; - - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id - const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id - const int64_t split_offset = global_sid * Trait::kTileDim; - if (global_pid >= params.num_compress) return; - - const auto plan = params.plan_c[global_pid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_output = static_cast(params.kv_output) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - const auto score_bias = static_cast(params.score_bias) + split_offset; - if (plan.is_invalid()) return; - - const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; - // Compact output: one row per compress plan, indexed by `global_pid`. - const auto kv_out = kv_output + global_pid * Trait::kHeadDim; - const auto kv_buf_0 = kv_buffer + plan.read_page_0 * Trait::kPageElementSize; - const auto kv_buf_1 = kv_buffer + plan.read_page_1 * Trait::kPageElementSize; - const bool need_overlap = plan.seq_len > 4; - PDLWaitPrimary(); - c4_forward(kv_buf_0, kv_buf_1, kv_src, kv_out, score_bias, need_overlap, plan.buffer_len); -} - -template -WRITE_KERNEL void write_c4_prefill(const __grid_constant__ Compress4PrefillParams params) { - using namespace device; - using Trait = C4Trait; - using StorageIn = AlignedVector; - - const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t global_wid = global_tid / kWarpThreads; // warp id - const uint32_t global_pid = global_wid / Trait::kNumSplit; // plan id - const uint32_t global_sid = global_wid % Trait::kNumSplit; // split id - // split the contiguous `kHeadDim * 4` into `kNumSplit` tiles - // each warp handles 1 contiguous tile (in contrast, decode handle the strided head_dim) - const int64_t split_offset = global_sid * (Trait::kTileDim * 4); - if (global_pid >= params.num_write) return; - - const auto plan = params.plan_w[global_pid]; - const auto kv_input = static_cast(params.kv_input) + split_offset; - const auto kv_buffer = static_cast(params.kv_buffer) + split_offset; - if (plan.is_invalid()) return; - - // each warp will handle a contiguous region - const auto kv_src = kv_input + plan.ragged_id * Trait::kElementSize; - const auto kv_buf = kv_buffer + plan.write_loc * Trait::kElementSize; - const auto gmem = tile::Memory::warp(); - - PDLWaitPrimary(); - StorageIn data[4]; -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - data[i] = gmem.load(kv_src, i); - } - PDLTriggerSecondary(); -#pragma unroll - for (int32_t i = 0; i < 4; ++i) { - gmem.store(kv_buf, data[i], i); - } -} - -template -struct FlashCompress4Kernel { - static constexpr auto decode_kernel = flash_c4_decode; - static constexpr auto prefill_c_kernel = flash_c4_prefill; - static constexpr auto prefill_w_kernel = write_c4_prefill; - static constexpr uint32_t kBlockSize = 128; - static constexpr uint32_t kTileDim = kTileElements * device::kWarpThreads; - static constexpr uint32_t kNumSplit = kHeadDim / kTileDim; - static constexpr uint32_t kWarpsPerBlock = kBlockSize / device::kWarpThreads; - using Trait = C4Trait; - - static void run_decode( - const tvm::ffi::TensorView kv_buffer, - const tvm::ffi::TensorView kv_input, - const tvm::ffi::TensorView kv_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_d_) { - using namespace host; - - auto N = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 4, Trait::kElementSize}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_buffer); - TensorMatcher({N, Trait::kElementSize}) // kv score input - .with_dtype() - .with_device(device_) - .verify(kv_input); - TensorMatcher({N, kHeadDim}) // kv compressed output - .with_dtype() - .with_device(device_) - .verify(kv_output); - TensorMatcher({8, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - - const auto plan_d = compress::verify_plan_d(plan_d_, N, device_); - const auto batch_size = static_cast(N.unwrap()); - const auto params = Compress4DecodeParams{ - .kv_buffer = kv_buffer.data_ptr(), - .kv_input = kv_input.data_ptr(), - .kv_output = kv_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_d = plan_d, - .batch_size = batch_size, - }; - const uint32_t num_blocks = div_ceil(batch_size * kNumSplit, kWarpsPerBlock); - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(decode_kernel, params); - } - - static void run_prefill( - const tvm::ffi::TensorView kv_buffer, - const tvm::ffi::TensorView kv_input, - const tvm::ffi::TensorView kv_output, - const tvm::ffi::TensorView ape, - const tvm::ffi::TensorView plan_c_, - const tvm::ffi::TensorView plan_w_) { - using namespace host; - - auto N = SymbolicSize{"num_q_tokens"}; - auto C = SymbolicSize{"num_c_plans"}; - auto W = SymbolicSize{"num_w_plans"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({-1, 4, Trait::kElementSize}) // kv score - .with_dtype() - .with_device(device_) - .verify(kv_buffer); - TensorMatcher({N, Trait::kElementSize}) // kv score input (ragged) - .with_dtype() - .with_device(device_) - .verify(kv_input); - TensorMatcher({C, kHeadDim}) // kv compressed output (compact) - .with_dtype() - .with_device(device_) - .verify(kv_output); - TensorMatcher({8, kHeadDim}) // ape - .with_dtype() - .with_device(device_) - .verify(ape); - const auto plan_c = compress::verify_plan_c(plan_c_, C, device_); - const auto plan_w = compress::verify_plan_w(plan_w_, W, device_); - const auto device = device_.unwrap(); - const auto num_q_tokens = static_cast(N.unwrap()); - const auto num_c = static_cast(C.unwrap()); - const auto num_w = static_cast(W.unwrap()); - const auto params = Compress4PrefillParams{ - .kv_buffer = kv_buffer.data_ptr(), - .kv_input = kv_input.data_ptr(), - .kv_output = kv_output.data_ptr(), - .score_bias = ape.data_ptr(), - .plan_c = plan_c, - .plan_w = plan_w, - .num_compress = num_c, - .num_write = num_w, - }; - RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w"); - if (const auto num_c_blocks = div_ceil(num_c * kNumSplit, kWarpsPerBlock)) { - LaunchKernel(num_c_blocks, kBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_c_kernel, params); - } - if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerBlock)) { - LaunchKernel(num_w_blocks, kBlockSize, device) // - .enable_pdl(kUsePDL)(prefill_w_kernel, params); - } - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh deleted file mode 100644 index 3e4aaaf5f0..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/c_plan.cuh +++ /dev/null @@ -1,839 +0,0 @@ -#include -#include -#include - -#include -#include - -#include - -#include -#include - -#include -#include - -namespace host::compress { - -constexpr auto kDLUInt8 = DLDataType{.code = kDLUInt, .bits = 8, .lanes = 1}; - -using PlanC = CompressPlan; -using PlanW = WritePlan; -using PlanD = DecodePlan; - -using RID_T = int64_t; -using R2T_T = int32_t; -using F2S_T = int64_t; -using IDX_T = int64_t; - -/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 -SGL_DEVICE __host__ PlanW pack_w(uint32_t ragged_id, uint32_t batch_id, int32_t seq_len) { - return {static_cast(ragged_id | batch_id << 16), seq_len}; -} - -/// NOTE: for the internal use, we pack the ragged and batch id, since both not exceed 65536 -SGL_DEVICE uint2 unpack_w(PlanW plan) { - return {static_cast(plan.ragged_id), static_cast(plan.ragged_id >> 16)}; -} - -struct Prefill0Params { - PlanC* plan_c; - PlanW* plan_w; - const IDX_T* seq_lens_ptr; // [batch_size] - const IDX_T* extend_lens_ptr; // [batch_size] - uint32_t batch_size; - uint32_t num_q_tokens; - int32_t compress_ratio; - int32_t swa_page_size; - int32_t mtp_pad; -}; - -struct Prefill1Params { - PlanC* plan_c; - PlanW* plan_w; - const RID_T* rid_ptr; // [batch_size] - const R2T_T* r2t_ptr; // [num_reqs, stride_r2t] - const F2S_T* f2s_ptr; // [num_swa_slots] - int64_t stride_r2t; - uint32_t num_c; - uint32_t num_w; - uint32_t num_c_padded; - uint32_t num_w_padded; - uint32_t num_work; - int32_t swa_page_size; - int32_t ring_size; - int32_t compress_ratio; -}; - -struct DecodeParams { - PlanD* plan_d; - const RID_T* rid_ptr; // [batch_size] - const R2T_T* r2t_ptr; // [num_reqs, stride_r2t] - const F2S_T* f2s_ptr; // [num_swa_slots] - const IDX_T* seq_ptr; // [batch_size] - int64_t stride_r2t; - uint32_t batch_size; - int32_t swa_page_size; - int32_t ring_size; - int32_t compress_ratio; -}; - -struct Prefill1ParamsLegacy { - PlanC* plan_c; - PlanW* plan_w; - const RID_T* rid_ptr; // [batch_size] - uint32_t num_c; - uint32_t num_w; - uint32_t num_c_padded; - uint32_t num_w_padded; - uint32_t num_work; - int32_t compress_ratio; -}; - -struct DecodeParamsLegacy { - PlanD* plan_d; - const RID_T* rid_ptr; // [batch_size] - const IDX_T* seq_ptr; // [batch_size] - uint32_t batch_size; - int32_t compress_ratio; -}; - -inline constexpr uint32_t kMaxPrefillBatchSize = 1024; - -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(device::kWarpThreads == 32); -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { -#ifndef USE_ROCM - uint32_t n = __shfl_up_sync(device::kFullMask, val, offset); -#else - uint32_t n = __shfl_up(val, offset, 32); -#endif - if (lane_id >= offset) val += n; - } - return val; -} - -/// Warp-wide max/min for integer types. `device::warp::reduce_max` routes through -/// `dtype_trait::max` which is only specialized for FP types. -SGL_DEVICE uint32_t warp_reduce_max_u32(uint32_t val) { -#pragma unroll - for (uint32_t mask = 16; mask > 0; mask >>= 1) { -#ifndef USE_ROCM - val = max(val, __shfl_xor_sync(device::kFullMask, val, mask, 32)); -#else - val = max(val, __shfl_xor(val, mask, 32)); -#endif - } - return val; -} - -SGL_DEVICE uint32_t warp_reduce_min_u32(uint32_t val) { -#pragma unroll - for (uint32_t mask = 16; mask > 0; mask >>= 1) { -#ifndef USE_ROCM - val = min(val, __shfl_xor_sync(device::kFullMask, val, mask, 32)); -#else - val = min(val, __shfl_xor(val, mask, 32)); -#endif - } - return val; -} - -__global__ __launch_bounds__(1024, 1) // - void plan_compress_prefill_kernel0(const Prefill0Params params) { - using namespace device; - const auto tx = threadIdx.x; - const auto block_size = kMaxPrefillBatchSize; - constexpr auto kNumWarps = kMaxPrefillBatchSize / kWarpThreads; - const auto cr = params.compress_ratio; - const auto sps = params.swa_page_size; - const bool is_overlap = (cr == 4); - const int32_t window_size = cr * (is_overlap ? 2 : 1); - - alignas(128) __shared__ uint32_t counter_c; - alignas(128) __shared__ uint32_t counter_w; - __shared__ int32_t s_seq_len[kMaxPrefillBatchSize]; - __shared__ int32_t s_prefix_len[kMaxPrefillBatchSize]; - __shared__ uint32_t warp_max[kNumWarps]; - __shared__ uint32_t warp_min[kNumWarps]; - __shared__ uint32_t s_max_extend; - __shared__ uint32_t s_min_extend; - - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - // === Stage A: load per-batch fields, init shared scratch === - int32_t seq_len = 0, extend_len = 0, prefix_len = 0; - if (tx < params.batch_size) { - seq_len = static_cast(params.seq_lens_ptr[tx]); - extend_len = static_cast(params.extend_lens_ptr[tx]); - prefix_len = seq_len - extend_len; - s_seq_len[tx] = seq_len; - s_prefix_len[tx] = prefix_len; - } - if (tx == 0) { - counter_c = 0; - counter_w = 0; - } - if (tx < kNumWarps) { - warp_max[tx] = 0; - warp_min[tx] = 0xFFFFFFFFu; - } - - // === Stage B: min/max(extend_len) for MTP-uniform detection === - // For min, treat threads outside `batch_size` as +inf so they don't pull the min down. - const uint32_t e_for_max = static_cast(extend_len); - const uint32_t e_for_min = (tx < params.batch_size) ? e_for_max : 0xFFFFFFFFu; - warp_max[warp_id] = warp_reduce_max_u32(e_for_max); - warp_min[warp_id] = warp_reduce_min_u32(e_for_min); - __syncthreads(); - if (warp_id == 0) { - s_max_extend = warp_reduce_max_u32(warp_max[lane_id]); - s_min_extend = warp_reduce_min_u32(warp_min[lane_id]); - } - __syncthreads(); - - const auto num_q = params.num_q_tokens; - // MTP-uniform: every batch shares the same small extend_len `E`, so we can decompose - // a global token id `k` into (batch_id, j) = (k / E, k % E) and skip the per-batch loop. - const bool is_mtp_extend = (s_min_extend == s_max_extend) && (s_max_extend > 0) && (s_max_extend <= 32); - - // === Stage C: emit valid plans, slot allocation via shared-mem atomicAdd === - if (is_mtp_extend) { - // Path 1: token-driven. Each global token id maps to exactly one (batch_id, j). - const uint32_t E = s_max_extend; - for (uint32_t k = tx; k < num_q; k += block_size) { - const uint32_t batch_id = k / E; - const uint32_t j = k % E; - const int32_t pl = s_prefix_len[batch_id]; - const int32_t sl = s_seq_len[batch_id]; - const int32_t position = pl + static_cast(j); - const uint32_t ragged_id = k; - - if ((position + 1) % cr == 0) { - const int32_t buffer_len = window_size - min(static_cast(j) + 1, window_size); - const uint32_t out_idx = atomicAdd(&counter_c, 1u); - params.plan_c[out_idx] = { - .seq_len = static_cast(position + 1), - .ragged_id = static_cast(ragged_id), - .buffer_len = static_cast(buffer_len), - .read_page_0 = -1, - .read_page_1 = static_cast(batch_id), - }; - } - - const int32_t last_c_pos = (sl / cr) * cr; - const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); - bool do_write = position >= first_w_pos; - if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); - if (do_write) { - const uint32_t out_idx = atomicAdd(&counter_w, 1u); - params.plan_w[out_idx] = pack_w(ragged_id, batch_id, position + 1); - } - } - } else { - // Path 2: general prefill (long extend_len). Iterate batches in an outer loop; - // the whole block sweeps each batch's tokens in parallel. - uint32_t base_e = 0; - for (uint32_t batch_id = 0; batch_id < params.batch_size; ++batch_id) { - const int32_t pl = s_prefix_len[batch_id]; - const int32_t sl = s_seq_len[batch_id]; - const int32_t el = sl - pl; - const int32_t last_c_pos = (sl / cr) * cr; - const int32_t first_w_pos = min(last_c_pos - (is_overlap ? cr : 0), sl - params.mtp_pad); - for (int32_t j = static_cast(tx); j < el; j += static_cast(block_size)) { - const int32_t position = pl + j; - const uint32_t ragged_id = base_e + static_cast(j); - - if ((position + 1) % cr == 0) { - const int32_t buffer_len = window_size - min(j + 1, window_size); - const uint32_t out_idx = atomicAdd(&counter_c, 1u); - params.plan_c[out_idx] = { - .seq_len = static_cast(position + 1), - .ragged_id = static_cast(ragged_id), - .buffer_len = static_cast(buffer_len), - .read_page_0 = -1, - .read_page_1 = static_cast(batch_id), - }; - } - - bool do_write = position >= first_w_pos; - if (!do_write && is_overlap) do_write = (position % sps) >= (sps - cr); - if (do_write) { - const uint32_t out_idx = atomicAdd(&counter_w, 1u); - params.plan_w[out_idx] = pack_w(ragged_id, static_cast(batch_id), position + 1); - } - } - base_e += static_cast(el); - } - } - __syncthreads(); - - // === Stage D: pad [counter_c, num_q) / [counter_w, num_q) with invalid === - const auto total_c = counter_c; - const auto total_w = counter_w; - for (uint32_t k = total_c + tx; k < num_q; k += block_size) { - params.plan_c[k] = PlanC::invalid(); - } - for (uint32_t k = total_w + tx; k < num_q; k += block_size) { - params.plan_w[k] = PlanW::invalid(); - } -} - -/// NOTE: stage 1 -__global__ void plan_compress_prefill_kernel_1(const Prefill1Params params) { - const auto idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= params.num_work) return; - auto plan_c = idx < params.num_c ? params.plan_c[idx] : PlanC::invalid(); - auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid(); - - const auto compute_loc = [&](int32_t swa_loc) { - const auto swa_page = swa_loc / params.swa_page_size; - const auto ring_offset = swa_loc % params.ring_size; - return swa_page * params.ring_size + ring_offset; - }; - - if (!plan_c.is_invalid()) { // 1. in bound. 2. not masked - if (plan_c.buffer_len > 0) { - const auto batch_id = plan_c.read_page_1; - const auto rid = params.rid_ptr[batch_id]; - const auto mapping = params.r2t_ptr + rid * params.stride_r2t; - // `seq_len` should be ratio-aligned here - const auto position_1 = static_cast(plan_c.seq_len - 1); - // only used for c4, harmless for c128 - const auto position_0 = max(position_1 - params.compress_ratio, 0); - const auto raw_loc_0 = mapping[position_0]; - const auto raw_loc_1 = mapping[position_1]; - const auto swa_loc_0 = params.f2s_ptr[raw_loc_0]; - const auto swa_loc_1 = params.f2s_ptr[raw_loc_1]; - plan_c.read_page_0 = compute_loc(swa_loc_0) / params.compress_ratio; - plan_c.read_page_1 = compute_loc(swa_loc_1) / params.compress_ratio; - params.plan_c[idx] = plan_c; - } - } else if (idx < params.num_c_padded) { - params.plan_c[idx] = PlanC::invalid(); - } - - if (!plan_w.is_invalid()) { // 1. in bound. 2. not masked - const auto [ragged_id, batch_id] = unpack_w(plan_w); - const auto rid = params.rid_ptr[batch_id]; - const auto mapping = params.r2t_ptr + rid * params.stride_r2t; - // `seq_len` (`write_loc`) may not be aligned here - const auto position = static_cast(plan_w.write_loc - 1); - const auto raw_loc = mapping[position]; - const auto swa_loc = params.f2s_ptr[raw_loc]; - plan_w.ragged_id = ragged_id; - plan_w.write_loc = compute_loc(swa_loc); - params.plan_w[idx] = plan_w; - } else if (idx < params.num_w_padded) { - params.plan_w[idx] = PlanW::invalid(); - } -} - -__global__ void plan_compress_decode_kernel(const DecodeParams params) { - const auto idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= params.batch_size) return; - const auto rid = params.rid_ptr[idx]; - const auto mapping = params.r2t_ptr + rid * params.stride_r2t; - const auto compute_loc = [&](int32_t swa_loc) { - const auto swa_page = swa_loc / params.swa_page_size; - const auto ring_offset = swa_loc % params.ring_size; - return swa_page * params.ring_size + ring_offset; - }; - const auto seq_len = static_cast(params.seq_ptr[idx]); - const auto position_1 = static_cast(seq_len - 1); - const auto position_0 = max(position_1 - params.compress_ratio, 0); - const auto raw_loc_0 = mapping[position_0]; - const auto raw_loc_1 = mapping[position_1]; - const auto swa_loc_0 = params.f2s_ptr[raw_loc_0]; - const auto swa_loc_1 = params.f2s_ptr[raw_loc_1]; - const auto write_loc = compute_loc(swa_loc_1); - const auto read_page_0 = compute_loc(swa_loc_0) / params.compress_ratio; - const auto read_page_1 = write_loc / params.compress_ratio; - params.plan_d[idx] = { - .seq_len = static_cast(seq_len), - .write_loc = write_loc, - .read_page_0 = read_page_0, - .read_page_1 = read_page_1, - }; -} - -__global__ void plan_compress_prefill_legacy_kernel(const Prefill1ParamsLegacy params) { - const auto idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= params.num_work) return; - auto plan_c = idx < params.num_c ? params.plan_c[idx] : PlanC::invalid(); - auto plan_w = idx < params.num_w ? params.plan_w[idx] : PlanW::invalid(); - - /// Per-request ring buffer slot translation: - /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 - /// - c128: page = rid; slot = rid * 128 + position % 128 - const auto legacy_compute_page = [&](int32_t rid, int32_t position) { - if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); - return rid; // c128 - }; - const auto legacy_compute_loc = [&](int32_t rid, int32_t position) { - const auto remainder = position % params.compress_ratio; - return legacy_compute_page(rid, position) * params.compress_ratio + remainder; - }; - - if (!plan_c.is_invalid()) { - const auto batch_id = plan_c.read_page_1; - const auto rid = static_cast(params.rid_ptr[batch_id]); - // `seq_len` is ratio-aligned for compress events - const auto position_1 = static_cast(plan_c.seq_len) - 1; - const auto position_0 = max(position_1 - params.compress_ratio, 0); - plan_c.read_page_0 = legacy_compute_page(rid, position_0); - plan_c.read_page_1 = legacy_compute_page(rid, position_1); - params.plan_c[idx] = plan_c; - } else if (idx < params.num_c_padded) { - params.plan_c[idx] = PlanC::invalid(); - } - - if (!plan_w.is_invalid()) { - const auto [ragged_id, batch_id] = unpack_w(plan_w); - const auto rid = static_cast(params.rid_ptr[batch_id]); - // `write_loc` carries (position + 1) at this stage; may not be ratio-aligned - const auto position = static_cast(plan_w.write_loc) - 1; - plan_w.ragged_id = ragged_id; - plan_w.write_loc = legacy_compute_loc(rid, position); - params.plan_w[idx] = plan_w; - } else if (idx < params.num_w_padded) { - params.plan_w[idx] = PlanW::invalid(); - } -} - -__global__ void plan_compress_decode_legacy_kernel(const DecodeParamsLegacy params) { - const auto idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= params.batch_size) return; - /// Per-request ring buffer slot translation: - /// - c4: page = rid * 2 + (position / 4) % 2; slot = page * 4 + position % 4 - /// - c128: page = rid; slot = rid * 128 + position % 128 - const auto legacy_compute_page = [&](int32_t rid, int32_t position) { - if (params.compress_ratio == 4) return rid * 2 + ((position / 4) & 1); - return rid; // c128 - }; - const auto legacy_compute_loc = [&](int32_t rid, int32_t position) { - const auto remainder = position % params.compress_ratio; - return legacy_compute_page(rid, position) * params.compress_ratio + remainder; - }; - const auto rid = static_cast(params.rid_ptr[idx]); - const auto seq_len = static_cast(params.seq_ptr[idx]); - const auto position_1 = seq_len - 1; - const auto position_0 = max(position_1 - params.compress_ratio, 0); - const auto write_loc = legacy_compute_loc(rid, position_1); - const auto read_page_0 = legacy_compute_page(rid, position_0); - const auto read_page_1 = legacy_compute_page(rid, position_1); - params.plan_d[idx] = { - .seq_len = static_cast(seq_len), - .write_loc = write_loc, - .read_page_0 = read_page_0, - .read_page_1 = read_page_1, - }; -} - -using PrefillPlan = tvm::ffi::Tuple; - -/** - * \brief Build c4/c128 prefill plan tensors. CPU-resident. - * Inputs (all CPU-resident): - * @param req_pool_indices `[batch_size]` int64_t - * @param req_to_token `[num_reqs, max_tokens_per_req]` int64_t - * @param full_to_swa `[num_swa_slots]` int64_t - * @param seq_lens `[batch_size]` int64 - * @param extend_lens `[batch_size]` int64 - * @param compress_plan `[num_q_tokens, 16]` uint8 (output) - * @param write_plan `[num_q_tokens, 8]` uint8 (output) - * @param compress_ratio 4 for c4, 128 for c128 - * @param use_cuda_graph Whether the plans will be used with cuda graph (affects padding) - * @return (compress plan tensor, write plan tensor) - */ -inline PrefillPlan plan_compress_prefill( - const tvm::ffi::TensorView req_pool_indices, // GPU - const tvm::ffi::TensorView req_to_token, // GPU - const tvm::ffi::TensorView full_to_swa, // GPU - const tvm::ffi::TensorView seq_lens, // CPU/GPU - const tvm::ffi::TensorView extend_lens, // CPU/GPU - const tvm::ffi::TensorView pin_buffer, // CPU - const uint32_t num_q_tokens, - const int32_t compress_ratio, - const int32_t swa_page_size, - const int32_t ring_size, - const bool use_cuda_graph) { - auto B = SymbolicSize{"batch_size"}; - auto N = SymbolicSize{"num_q_tokens"}; - auto cpu_or_gpu = SymbolicDevice{}; - auto device_ = SymbolicDevice{}; - cpu_or_gpu.set_options(); - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({-1, -1}) // - .with_dtype() - .with_device(device_) - .verify(req_to_token); - TensorMatcher({-1}) // - .with_dtype() - .with_device(device_) - .verify(full_to_swa); - TensorMatcher({B}) // - .with_dtype() - .with_device(cpu_or_gpu) - .verify(seq_lens) - .verify(extend_lens); - TensorMatcher({-1}) // - .with_dtype() - .with_device() - .verify(pin_buffer); - - const bool is_overlap = (compress_ratio == 4); - const int32_t window_size = compress_ratio * (is_overlap ? 2 : 1); - - const auto seq_ptr = static_cast(seq_lens.data_ptr()); - const auto ext_ptr = static_cast(extend_lens.data_ptr()); - const auto rid_ptr = static_cast(req_pool_indices.data_ptr()); - const auto r2t_ptr = static_cast(req_to_token.data_ptr()); - const auto f2s_ptr = static_cast(full_to_swa.data_ptr()); - - const auto batch_size = static_cast(B.unwrap()); - constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); - RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); - RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); - // `swa_page_size` >= `ring_size` >= `compress_ratio` - RuntimeCheck(swa_page_size % ring_size == 0 && ring_size % compress_ratio == 0); - - const auto device = device_.unwrap(); - const auto stream = LaunchKernel::resolve_device(device); - - constexpr int32_t kMaxMTPDraftTokens = 4; - const auto mtp_pad = std::min(ring_size - compress_ratio, kMaxMTPDraftTokens); - - if (cpu_or_gpu.unwrap().device_type == kDLGPU) { - // GPU input path: kernel0 builds the (CPU-loop-equivalent) plan metadata directly - // on device, padding to num_q_tokens with invalid; kernel_1 then finalizes the - // SWA-translated read/write locations. Used for MTP / cuda-graph capture where - // a host sync would be expensive. - RuntimeCheck(batch_size <= kMaxPrefillBatchSize, "GPU plan only support batch size up to ", kMaxPrefillBatchSize); - auto C = ffi::empty({num_q_tokens, sizeof(PlanC)}, kDLUInt8, device); - auto W = ffi::empty({num_q_tokens, sizeof(PlanW)}, kDLUInt8, device); - const auto params0 = Prefill0Params{ - .plan_c = static_cast(C.data_ptr()), - .plan_w = static_cast(W.data_ptr()), - .seq_lens_ptr = seq_ptr, - .extend_lens_ptr = ext_ptr, - .batch_size = batch_size, - .num_q_tokens = num_q_tokens, - .compress_ratio = compress_ratio, - .swa_page_size = swa_page_size, - .mtp_pad = mtp_pad, - }; - LaunchKernel(1, kMaxPrefillBatchSize, device)(plan_compress_prefill_kernel0, params0); - // kernel_1 sees the already-padded buffers, so num_c == num_w == num_padded == num_q_tokens. - const auto params1 = Prefill1Params{ - .plan_c = static_cast(C.data_ptr()), - .plan_w = static_cast(W.data_ptr()), - .rid_ptr = rid_ptr, - .r2t_ptr = r2t_ptr, - .f2s_ptr = f2s_ptr, - .stride_r2t = req_to_token.stride(0), - .num_c = num_q_tokens, - .num_w = num_q_tokens, - .num_c_padded = num_q_tokens, - .num_w_padded = num_q_tokens, - .num_work = num_q_tokens, - .swa_page_size = swa_page_size, - .ring_size = ring_size, - .compress_ratio = compress_ratio, - }; - const auto block_size_1 = 256; - const auto num_blocks_1 = div_ceil(params1.num_work, block_size_1); - LaunchKernel(num_blocks_1, block_size_1, device)(plan_compress_prefill_kernel_1, params1); - return PrefillPlan{std::move(C), std::move(W)}; - } - - // CPU input path: only here do we need the pinned scratch buffer. - const auto pin_buffer_bytes = static_cast(pin_buffer.numel()) * sizeof(uint8_t); - RuntimeCheck(pin_buffer_bytes >= num_q_tokens * (sizeof(PlanC) + sizeof(PlanW))); - const auto plan_c_ptr = reinterpret_cast(pin_buffer.data_ptr()); - const auto plan_w_ptr = reinterpret_cast(plan_c_ptr + num_q_tokens); - - uint32_t counter = 0; - uint32_t counter_c = 0; - uint32_t counter_w = 0; - - const auto should_compress = [=](int32_t position) { return (position + 1) % compress_ratio == 0; }; - for (const auto i : irange(batch_size)) { - const int32_t seq_len = seq_ptr[i]; - const int32_t extend_len = ext_ptr[i]; - const int32_t prefix_len = seq_len - extend_len; - const int32_t last_c_pos = seq_len / compress_ratio * compress_ratio; - const int32_t first_w_pos = last_c_pos - (is_overlap ? compress_ratio : 0); - RuntimeCheck(0 < extend_len && extend_len <= seq_len); - const auto should_write = [=](int32_t position) { - if (position >= first_w_pos) return true; - return is_overlap && position % swa_page_size >= (swa_page_size - compress_ratio); - }; - for (const auto j : irange(extend_len)) { - const int32_t position = prefix_len + j; - const int32_t ragged_id = counter + j; - if (should_compress(position)) { - const auto buffer_len = window_size - std::min(j + 1, window_size); - plan_c_ptr[counter_c++] = { - .seq_len = static_cast(position + 1), - .ragged_id = static_cast(ragged_id), - .buffer_len = static_cast(buffer_len), - // to be filled by kernel - .read_page_0 = -1, - .read_page_1 = static_cast(i), - }; - } - if (should_write(position)) { - plan_w_ptr[counter_w++] = pack_w(ragged_id, i, position + 1); - } - } - counter += extend_len; - } - RuntimeCheck(counter == num_q_tokens); - - const auto copy_to_device = [stream](void* cuda_ptr, auto* host_ptr, size_t count) { - const auto size_bytes = count * sizeof(*host_ptr); - RuntimeDeviceCheck(cudaMemcpyAsync(cuda_ptr, host_ptr, size_bytes, cudaMemcpyHostToDevice, stream)); - }; - const auto num_c_padded = use_cuda_graph ? num_q_tokens : counter_c; - const auto num_w_padded = use_cuda_graph ? num_q_tokens : counter_w; - auto C = ffi::empty({num_c_padded, sizeof(PlanC)}, kDLUInt8, device); - auto W = ffi::empty({num_w_padded, sizeof(PlanW)}, kDLUInt8, device); - copy_to_device(C.data_ptr(), plan_c_ptr, counter_c); - copy_to_device(W.data_ptr(), plan_w_ptr, counter_w); - const auto params = Prefill1Params{ - .plan_c = static_cast(C.data_ptr()), - .plan_w = static_cast(W.data_ptr()), - .rid_ptr = rid_ptr, - .r2t_ptr = r2t_ptr, - .f2s_ptr = f2s_ptr, - .stride_r2t = req_to_token.size(1), - .num_c = counter_c, - .num_w = counter_w, - .num_c_padded = num_c_padded, - .num_w_padded = num_w_padded, - .num_work = std::max(num_c_padded, num_w_padded), - .swa_page_size = swa_page_size, - .ring_size = ring_size, - .compress_ratio = compress_ratio, - }; - const auto block_size = 256; - const auto num_blocks = div_ceil(params.num_work, block_size); - LaunchKernel(num_blocks, block_size, device)(plan_compress_prefill_kernel_1, params); - return PrefillPlan{std::move(C), std::move(W)}; -} - -inline tvm::ffi::Tensor plan_compress_decode( - const tvm::ffi::TensorView req_pool_indices, // GPU - const tvm::ffi::TensorView req_to_token, // GPU - const tvm::ffi::TensorView full_to_swa, // GPU - const tvm::ffi::TensorView seq_lens, // CPU/GPU - const int32_t compress_ratio, - const int32_t swa_page_size, - const int32_t ring_size) { - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({-1, -1}) // - .with_dtype() - .with_device(device_) - .verify(req_to_token); - TensorMatcher({-1}) // - .with_dtype() - .with_device(device_) - .verify(full_to_swa); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(seq_lens); - - const auto batch_size = static_cast(B.unwrap()); - const auto device = device_.unwrap(); - auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device); - const auto params = DecodeParams{ - .plan_d = static_cast(D.data_ptr()), - .rid_ptr = static_cast(req_pool_indices.data_ptr()), - .r2t_ptr = static_cast(req_to_token.data_ptr()), - .f2s_ptr = static_cast(full_to_swa.data_ptr()), - .seq_ptr = static_cast(seq_lens.data_ptr()), - .stride_r2t = req_to_token.size(1), - .batch_size = batch_size, - .swa_page_size = swa_page_size, - .ring_size = ring_size, - .compress_ratio = compress_ratio, - }; - const auto block_size = 256; - const auto num_blocks = div_ceil(batch_size, block_size); - LaunchKernel(num_blocks, block_size, device)(plan_compress_decode_kernel, params); - return D; -} - -/** - * \brief Build c4/c128 prefill plan tensors for the legacy non-paged ring - * buffer. Uses only `req_pool_indices` to derive ring slots: - * - c4 (overlap): each request occupies 2 contiguous pages (8 token slots) - * - c128: each request occupies 1 page (128 token slots) - * - * Inputs: - * @param req_pool_indices `[batch_size]` int64 (GPU) - * @param seq_lens `[batch_size]` int64 (CPU) - * @param extend_lens `[batch_size]` int64 (CPU) - * @param pin_buffer pinned scratch (CPU uint8) - * @return (compress plan tensor, write plan tensor) - */ -inline PrefillPlan plan_compress_prefill_legacy( - const tvm::ffi::TensorView req_pool_indices, // GPU - const tvm::ffi::TensorView seq_lens, // CPU - const tvm::ffi::TensorView extend_lens, // CPU - const tvm::ffi::TensorView pin_buffer, // CPU - const uint32_t num_q_tokens, - const int32_t compress_ratio, - const bool use_cuda_graph) { - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({B}) // - .with_dtype() - .with_device() - .verify(seq_lens) - .verify(extend_lens); - TensorMatcher({-1}) // - .with_dtype() - .with_device() - .verify(pin_buffer); - - const auto pin_buffer_bytes = static_cast(pin_buffer.numel()) * sizeof(uint8_t); - RuntimeCheck(pin_buffer_bytes >= num_q_tokens * (sizeof(PlanC) + sizeof(PlanW))); - const auto plan_c_ptr = reinterpret_cast(pin_buffer.data_ptr()); - const auto plan_w_ptr = reinterpret_cast(plan_c_ptr + num_q_tokens); - - const bool is_overlap = (compress_ratio == 4); - const auto seq_ptr = static_cast(seq_lens.data_ptr()); - const auto ext_ptr = static_cast(extend_lens.data_ptr()); - const auto rid_ptr = static_cast(req_pool_indices.data_ptr()); - - const auto window_size = compress_ratio * (is_overlap ? 2 : 1); - const auto batch_size = static_cast(B.unwrap()); - constexpr auto kMaxTokens = static_cast(std::numeric_limits::max()); - RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); - RuntimeCheck(batch_size <= num_q_tokens && num_q_tokens <= kMaxTokens); - - uint32_t counter = 0; - uint32_t counter_c = 0; - uint32_t counter_w = 0; - const auto should_compress = [=](int32_t position) { return (position + 1) % compress_ratio == 0; }; - for (const auto i : irange(batch_size)) { - const int32_t seq_len = seq_ptr[i]; - const int32_t extend_len = ext_ptr[i]; - const int32_t prefix_len = seq_len - extend_len; - const int32_t last_c_pos = seq_len / compress_ratio * compress_ratio; - const int32_t first_w_pos = last_c_pos - (is_overlap ? compress_ratio : 0); - RuntimeCheck(0 < extend_len && extend_len <= seq_len); - const auto should_write = [=](int32_t position) { return position >= first_w_pos; }; - for (const auto j : irange(extend_len)) { - const int32_t position = prefix_len + j; - const int32_t ragged_id = counter + j; - if (should_compress(position)) { - const auto buffer_len = window_size - std::min(j + 1, window_size); - plan_c_ptr[counter_c++] = { - .seq_len = static_cast(position + 1), - .ragged_id = static_cast(ragged_id), - .buffer_len = static_cast(buffer_len), - // to be filled by kernel - .read_page_0 = -1, - .read_page_1 = static_cast(i), - }; - } - if (should_write(position)) { - plan_w_ptr[counter_w++] = pack_w(ragged_id, i, position + 1); - } - } - counter += extend_len; - } - RuntimeCheck(counter == num_q_tokens); - - const auto device = device_.unwrap(); - const auto stream = LaunchKernel::resolve_device(device); - const auto copy_to_device = [stream](void* cuda_ptr, auto* host_ptr, size_t count) { - const auto size_bytes = count * sizeof(*host_ptr); - RuntimeDeviceCheck(cudaMemcpyAsync(cuda_ptr, host_ptr, size_bytes, cudaMemcpyHostToDevice, stream)); - }; - const auto num_c_padded = use_cuda_graph ? num_q_tokens : counter_c; - const auto num_w_padded = use_cuda_graph ? num_q_tokens : counter_w; - auto C = ffi::empty({num_c_padded, sizeof(PlanC)}, kDLUInt8, device); - auto W = ffi::empty({num_w_padded, sizeof(PlanW)}, kDLUInt8, device); - copy_to_device(C.data_ptr(), plan_c_ptr, counter_c); - copy_to_device(W.data_ptr(), plan_w_ptr, counter_w); - const auto params = Prefill1ParamsLegacy{ - .plan_c = static_cast(C.data_ptr()), - .plan_w = static_cast(W.data_ptr()), - .rid_ptr = rid_ptr, - .num_c = counter_c, - .num_w = counter_w, - .num_c_padded = num_c_padded, - .num_w_padded = num_w_padded, - .num_work = std::max(num_c_padded, num_w_padded), - .compress_ratio = compress_ratio, - }; - const auto block_size = 256; - const auto num_blocks = div_ceil(params.num_work, block_size); - if (num_blocks > 0) { - LaunchKernel(num_blocks, block_size, device)(plan_compress_prefill_legacy_kernel, params); - } - return PrefillPlan{std::move(C), std::move(W)}; -} - -inline tvm::ffi::Tensor plan_compress_decode_legacy( - const tvm::ffi::TensorView req_pool_indices, // GPU - const tvm::ffi::TensorView seq_lens, // GPU - const int32_t compress_ratio) { - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(req_pool_indices); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(seq_lens); - RuntimeCheck(compress_ratio == 4 || compress_ratio == 128); - - const auto batch_size = static_cast(B.unwrap()); - const auto device = device_.unwrap(); - auto D = ffi::empty({batch_size, sizeof(PlanD)}, kDLUInt8, device); - const auto params = DecodeParamsLegacy{ - .plan_d = static_cast(D.data_ptr()), - .rid_ptr = static_cast(req_pool_indices.data_ptr()), - .seq_ptr = static_cast(seq_lens.data_ptr()), - .batch_size = batch_size, - .compress_ratio = compress_ratio, - }; - const auto block_size = 256; - const auto num_blocks = div_ceil(batch_size, block_size); - LaunchKernel(num_blocks, block_size, device)(plan_compress_decode_legacy_kernel, params); - return D; -} - -} // namespace host::compress - -using namespace host::compress; // expose binding diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh deleted file mode 100644 index 46acaa9c46..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/common.cuh +++ /dev/null @@ -1,208 +0,0 @@ -#include -#include - -#include - -#include - -namespace host::compress { - -using PlanResult = tvm::ffi::Tuple; - -struct CompressParams { - PrefillPlan* __restrict__ compress_plan; - PrefillPlan* __restrict__ write_plan; - const int64_t* __restrict__ seq_lens; - const int64_t* __restrict__ extend_lens; - uint32_t batch_size; - uint32_t num_tokens; - uint32_t compress_ratio; - bool is_overlap; -}; - -inline constexpr uint32_t kBlockSize = 1024; - -#define PLAN_KERNEL __global__ __launch_bounds__(kBlockSize, 1) inline - -PLAN_KERNEL void plan_prefill_cuda(const __grid_constant__ CompressParams params) { - const auto &[ - compress_plan, write_plan, seq_lens, extend_lens, // pointers - batch_size, num_tokens, compress_ratio, is_overlap // values - ] = params; - - __shared__ uint32_t compress_counter; - __shared__ uint32_t write_counter; - - uint32_t batch_id = 0; - uint32_t counter = 0; - uint32_t extend_len = extend_lens[0]; - - const auto tid = threadIdx.x; - if (tid == 0) { - compress_counter = 0; - write_counter = 0; - } - __syncthreads(); - - for (uint32_t i = tid; i < num_tokens; i += blockDim.x) { - const uint32_t ragged_id = i; - uint32_t j = ragged_id - counter; - while (j >= extend_len) { - j -= extend_len; - batch_id += 1; - if (batch_id >= batch_size) [[unlikely]] - break; - counter += extend_len; - extend_len = extend_lens[batch_id]; - } - if (batch_id >= batch_size) [[unlikely]] - break; - const uint32_t seq_len = seq_lens[batch_id]; - const uint32_t extend_len = extend_lens[batch_id]; - const uint32_t prefix_len = seq_len - extend_len; - const uint32_t ratio = compress_ratio * (1 + is_overlap); - const uint32_t window_len = j + 1 < ratio ? ratio - (j + 1) : 0; - const uint32_t position = prefix_len + j; - const auto plan = PrefillPlan{ - .ragged_id = ragged_id, - .batch_id = batch_id, - .position = position, - .window_len = window_len, - }; - const uint32_t start_write_pos = [seq_len, compress_ratio, is_overlap] { - const uint32_t pos = seq_len / compress_ratio * compress_ratio; - if (!is_overlap) return pos; - return pos >= compress_ratio ? pos - compress_ratio : 0; - }(); - if ((position + 1) % compress_ratio == 0) { - const auto write_pos = atomicAdd(&compress_counter, 1); - compress_plan[write_pos] = plan; - } - if (position >= start_write_pos) { - const auto write_pos = atomicAdd(&write_counter, 1); - write_plan[write_pos] = plan; - } - } - __syncthreads(); - constexpr auto kInvalid = static_cast(-1); - const auto kInvalidPlan = PrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; - const auto compress_count = compress_counter; - const auto write_count = write_counter; - for (uint32_t i = compress_count + tid; i < num_tokens; i += blockDim.x) { - compress_plan[i] = kInvalidPlan; - } - for (uint32_t i = write_count + tid; i < num_tokens; i += blockDim.x) { - write_plan[i] = kInvalidPlan; - } -} - -inline PlanResult plan_prefill_host(const CompressParams& params, const bool use_cuda_graph) { - const auto &[ - compress_ptr, write_ptr, seq_lens_ptr, extend_lens_ptr, // pointers - batch_size, num_tokens, compress_ratio, is_overlap // values - ] = params; - - uint32_t counter = 0; - uint32_t compress_counter = 0; - uint32_t write_counter = 0; - const auto ratio = compress_ratio * (1 + is_overlap); - for (const auto i : irange(batch_size)) { - const uint32_t seq_len = seq_lens_ptr[i]; - const uint32_t extend_len = extend_lens_ptr[i]; - const uint32_t prefix_len = seq_len - extend_len; - RuntimeCheck(0 < extend_len && extend_len <= seq_len); - /// NOTE: `start_write_pos` must be a multiple of `compress_ratio` - const uint32_t start_write_pos = [seq_len, compress_ratio, is_overlap] { - const uint32_t pos = seq_len / compress_ratio * compress_ratio; - if (!is_overlap) return pos; - /// NOTE: to avoid unsigned integer underflow, don't use `pos - compress_ratio` - return pos >= compress_ratio ? pos - compress_ratio : 0; - }(); - /// NOTE: `position` is within [prefix_len, seq_len) - for (const auto j : irange(extend_len)) { - const uint32_t position = prefix_len + j; - const auto plan = PrefillPlan{ - .ragged_id = counter + j, - .batch_id = i, - .position = position, - .window_len = ratio - std::min(j + 1, ratio), - }; - RuntimeCheck(plan.is_valid(compress_ratio, is_overlap), "Internal error!"); - if ((position + 1) % compress_ratio == 0) { - compress_ptr[compress_counter++] = plan; - } - if (position >= start_write_pos) { - write_ptr[write_counter++] = plan; - } - } - counter += extend_len; - } - RuntimeCheck(counter == num_tokens, "input size ", counter, " != num_q_tokens ", num_tokens); - if (!use_cuda_graph) return PlanResult{compress_counter, write_counter}; - constexpr auto kInvalid = static_cast(-1); - constexpr auto kInvalidPlan = PrefillPlan{kInvalid, kInvalid, kInvalid, kInvalid}; - for (const auto i : irange(compress_counter, num_tokens)) { - compress_ptr[i] = kInvalidPlan; - } - for (const auto i : irange(write_counter, num_tokens)) { - write_ptr[i] = kInvalidPlan; - } - return PlanResult{num_tokens, num_tokens}; -} - -inline PlanResult plan_prefill( - const tvm::ffi::TensorView extend_lens, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView compress_plan, - const tvm::ffi::TensorView write_plan, - const uint32_t compress_ratio, - const bool is_overlap, // for overlap transform, we have to keep 1 more extra window - const bool use_cuda_graph) { - auto N = SymbolicSize{"batch_size"}; - auto M = SymbolicSize{"num_tokens"}; - auto device = SymbolicDevice{}; - const bool is_cuda = [&] { - if (extend_lens.device().device_type == kDLCUDA) { - device.set_options(); - return true; - } else { - device.set_options(); - return false; - } - }(); - TensorMatcher({N}) // extend_lens and seq_lens - .with_dtype() - .with_device(device) - .verify(extend_lens) - .verify(seq_lens); - TensorMatcher({M, kPrefillPlanDim}) // compress_plan and write_plan - .with_dtype() - .with_device(device) - .verify(compress_plan) - .verify(write_plan); - - const auto params = CompressParams{ - .compress_plan = static_cast(compress_plan.data_ptr()), - .write_plan = static_cast(write_plan.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .extend_lens = static_cast(extend_lens.data_ptr()), - .batch_size = static_cast(N.unwrap()), - .num_tokens = static_cast(M.unwrap()), - .compress_ratio = compress_ratio, - .is_overlap = is_overlap, - }; - - if (!is_cuda) return plan_prefill_host(params, use_cuda_graph); - /// NOTE: cuda kernel plan is naturally compatible with cuda graph - LaunchKernel(1, kBlockSize, device.unwrap())(plan_prefill_cuda, params); - return PlanResult{params.num_tokens, params.num_tokens}; -} - -} // namespace host::compress - -namespace { - -[[maybe_unused]] -constexpr auto& plan_compress_prefill = host::compress::plan_prefill; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh deleted file mode 100644 index d3953578b9..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope.cuh +++ /dev/null @@ -1,254 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include - -#include - -#include -#include - -namespace { - -using Plan = device::compress::PrefillPlan; - -/// \brief common block size for memory-bound kernel -constexpr uint32_t kBlockSize = 128; -constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; - -struct FusedNormRopeParams { - void* __restrict__ input; - const void* __restrict__ weight; - float eps; - uint32_t num_works; - const void* __restrict__ handle; - const float* __restrict__ freqs_cis; - uint32_t compress_ratio; -}; - -enum class ForwardMode { - CompressExtend = 0, - CompressDecode = 1, - DefaultForward = 2, -}; - -template -__global__ void fused_norm_rope(const __grid_constant__ FusedNormRopeParams params) { - using namespace device; - using enum ForwardMode; - - constexpr int64_t kMaxVecSize = 16 / sizeof(DType); - constexpr int64_t kVecSize = std::min(kMaxVecSize, kHeadDim / kWarpThreads); - constexpr int64_t kLocalSize = kHeadDim / (kWarpThreads * kVecSize); - constexpr int64_t kRopeVecSize = kRopeDim / (kWarpThreads * 2); - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; - static_assert(kHeadDim % (kWarpThreads * kVecSize) == 0); - static_assert(kLocalSize * kVecSize * kWarpThreads == kHeadDim); - static_assert(kRopeDim % (kWarpThreads * 2) == 0); - static_assert(kRopeDim % (kVecSize * kLocalSize) == 0); - static_assert(kRopeSize <= kWarpThreads); - static_assert(kRopeVecSize == 1, "only support rope dim = 64"); - - const auto& [ - _input, _weight, eps, num_works, // norm - handle, freqs_cis, compress_ratio // rope - ] = params; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kNumWarps + warp_id; - - if (work_id >= num_works) return; - - DType* input; - int32_t position; - if constexpr (kMode == CompressExtend) { - const auto plan = static_cast(handle)[work_id]; - input = static_cast(_input) + plan.ragged_id * kHeadDim; - position = plan.position + 1 - compress_ratio; - if (plan.ragged_id == 0xFFFFFFFF) [[unlikely]] - return; - } else if constexpr (kMode == CompressDecode) { - input = static_cast(_input) + work_id * kHeadDim; - const auto seq_len = static_cast(handle)[work_id]; - if (seq_len % compress_ratio != 0) return; - position = seq_len - compress_ratio; - } else if constexpr (kMode == DefaultForward) { - input = static_cast(_input) + work_id * kHeadDim; - position = static_cast(handle)[work_id]; - } else { - static_assert(host::dependent_false_v, "Unsupported Mode"); - } - - using Storage = AlignedVector; - __shared__ Storage s_rope_input[kNumWarps][kRopeSize]; - - // prefetch freq - const auto mem_freq = tile::Memory::warp(); - const auto freq = mem_freq.load(freqs_cis + position * kRopeDim); - - PDLWaitPrimary(); - - // part 1: norm - { - const auto gmem = tile::Memory::warp(); - Storage input_vec[kLocalSize]; - Storage weight_vec[kLocalSize]; -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { - input_vec[i] = gmem.load(input, i); - } - -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { - weight_vec[i] = gmem.load(_weight, i); - } - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { -#pragma unroll - for (int j = 0; j < kVecSize; ++j) { - const auto fp32_input = cast(input_vec[i][j]); - sum_of_squares += fp32_input * fp32_input; - } - } - - sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + eps); - -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { -#pragma unroll - for (int j = 0; j < kVecSize; ++j) { - const auto fp32_input = cast(input_vec[i][j]); - const auto fp32_weight = cast(weight_vec[i][j]); - input_vec[i][j] = cast(fp32_input * norm_factor * fp32_weight); - } - } - - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; - -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { - if (i == kLocalSize - 1 && is_rope_lane) { - const auto rope_id = lane_id - (kWarpThreads - kRopeSize); - s_rope_input[warp_id][rope_id] = input_vec[i]; - } else { - gmem.store(input, input_vec[i], i); - } - } - - __syncwarp(); - } - - // part 2: rope - { - // mem elem = DType x 2 - using DTypex2_t = packed_t; - const auto mem_elem = tile::Memory::warp(); - const auto elem = mem_elem.load(s_rope_input[warp_id]); - const auto [x_real, x_imag] = cast(elem); - const auto [freq_real, freq_imag] = freq; - const fp32x2_t output = { - x_real * freq_real - x_imag * freq_imag, - x_real * freq_imag + x_imag * freq_real, - }; - mem_elem.store(input + (kHeadDim - kRopeDim), cast(output)); - } - - PDLTriggerSecondary(); -} - -template -struct FusedNormRopeKernel { - template - static constexpr auto fused_kernel = fused_norm_rope; - - static void forward( - const tvm::ffi::TensorView input, - const tvm::ffi::TensorView weight, - const tvm::ffi::TensorView handle, - const tvm::ffi::TensorView freqs_cis, - int32_t _mode, - float eps, - uint32_t compress_ratio) { - using namespace host; - using enum ForwardMode; - - const auto mode = static_cast(_mode); - - auto B = SymbolicSize{"num_q_tokens"}; - auto N = SymbolicSize{"num_compress_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, kHeadDim}) // input - .with_dtype() - .with_device(device_) - .verify(input); - TensorMatcher({kHeadDim}) // weight - .with_dtype() - .with_device(device_) - .verify(weight); - TensorMatcher({-1, kRopeDim}) // freqs_cis - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - switch (mode) { - case CompressExtend: - TensorMatcher({N, compress::kPrefillPlanDim}) // plan - .with_dtype() - .with_device(device_) - .verify(handle); - RuntimeCheck(compress_ratio > 0); - break; - case CompressDecode: - TensorMatcher({N}) // seq_len - .with_dtype() - .with_device(device_) - .verify(handle); - RuntimeCheck(compress_ratio > 0); - break; - case DefaultForward: - TensorMatcher({N}) // position - .with_dtype() - .with_device(device_) - .verify(handle); - RuntimeCheck(compress_ratio == 0); - break; - default: - Panic("unsupported forward mode: ", static_cast(mode)); - } - - // launch kernel - const auto num_compress_tokens = static_cast(N.unwrap()); - if (num_compress_tokens == 0) return; - const auto params = FusedNormRopeParams{ - .input = input.data_ptr(), - .weight = weight.data_ptr(), - .eps = eps, - .num_works = num_compress_tokens, - .handle = handle.data_ptr(), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .compress_ratio = compress_ratio, - }; - const auto num_blocks = div_ceil(num_compress_tokens, kNumWarps); - using KernelType = std::decay_t)>; - static constexpr KernelType kernel_table[3] = { - [static_cast(CompressExtend)] = fused_kernel, - [static_cast(CompressDecode)] = fused_kernel, - [static_cast(DefaultForward)] = fused_kernel, - }; - const auto kernel = kernel_table[static_cast(mode)]; - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh deleted file mode 100644 index a9cac17544..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/fused_norm_rope_v2.cuh +++ /dev/null @@ -1,643 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include -#include - -#include - -#include - -namespace { - -using PlanC = device::compress::CompressPlan; -using PlanD = device::compress::DecodePlan; -using deepseek_v4::fp8::cast_to_ue8m0; -using deepseek_v4::fp8::inv_scale_ue8m0; -using deepseek_v4::fp8::pack_fp8; - -SGL_DEVICE uint8_t quant_fp4_e2m1(float x) { - const float ax = fminf(fabsf(x), 6.0f); - uint8_t idx = 0; - idx += ax > 0.25f; - idx += ax > 0.75f; - idx += ax > 1.25f; - idx += ax > 1.75f; - idx += ax > 2.5f; - idx += ax > 3.5f; - idx += ax > 5.0f; - if (x < 0.0f && idx != 0) idx |= 0x8; - return idx; -} - -constexpr uint32_t kBlockSize = 256; -constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; - -struct FusedNormRopeStoreParams { - void* __restrict__ input; - const void* __restrict__ handle; // plan decode / compress - const void* __restrict__ weight; - const float* __restrict__ freqs_cis; - const int32_t* __restrict__ out_loc; - uint8_t* __restrict__ kvcache; - float eps; - uint32_t compress_ratio; - uint32_t num_tokens; -}; - -enum class ForwardMode : bool { - CompressExtend = 0, - CompressDecode = 1, -}; - -#define INDEXER_KERNEL __global__ __launch_bounds__(kBlockSize, 8) -#define FLASHMLA_KERNEL __global__ __launch_bounds__(kBlockSize, 8) - -// ---------------------------------------------------------------------------- -// Indexer variant: kHeadDim = 128, 1 token per *warp* (8 tokens per block). -// Each warp's 32 lanes cover the full 128-elem head_dim (kVecSize = 4 each). -// Cache layout: 132 bytes/token (128 fp8 nope + 4 fp32 scale). -// ---------------------------------------------------------------------------- -template -INDEXER_KERNEL void fused_norm_rope_indexer(const __grid_constant__ FusedNormRopeStoreParams params) { - using namespace device; - using enum ForwardMode; - - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kVecSize = 4; - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; - constexpr int64_t kPageBytes = 132ll << kPageBits; - static_assert(kHeadDim == kWarpThreads * kVecSize); - static_assert(kRopeDim == kWarpThreads * 2); - static_assert(kRopeSize <= kWarpThreads); - using Storage = AlignedVector; - using Float4 = AlignedVector; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kNumWarps + warp_id; - // Lanes whose 4-elem pack lies in the rope tail (= last `kRopeSize` packs). - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; - - if (work_id >= params.num_tokens) return; - - const auto input = static_cast(params.input) + work_id * kHeadDim; - int32_t position; - int32_t out_loc; - if constexpr (kMode == CompressExtend) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.is_invalid()) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[plan.ragged_id]; - } else if constexpr (kMode == CompressDecode) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.seq_len % params.compress_ratio != 0) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[work_id]; - } else { - static_assert(host::dependent_false_v, "Unsupported Mode"); - } - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - PDLWaitPrimary(); - Float4 data, freq; - - // part 1: norm - { - Storage input_vec, weight_vec; - input_vec.load(input, lane_id); - weight_vec.load(params.weight, lane_id); - if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - sum_of_squares += fp32_input * fp32_input; - } - - sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); - -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - const auto fp32_weight = cast(weight_vec[i]); - data[i] = fp32_input * norm_factor * fp32_weight; - } - } - - // part 2: rope (rope-lane only, 4 elems per lane = 2 (real, imag) pairs) - if (is_rope_lane) { - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto y_real = data[2]; - const auto y_imag = data[3]; - const auto freq_x_real = freq[0]; - const auto freq_x_imag = freq[1]; - const auto freq_y_real = freq[2]; - const auto freq_y_imag = freq[3]; - data[0] = x_real * freq_x_real - x_imag * freq_x_imag; - data[1] = x_real * freq_x_imag + x_imag * freq_x_real; - data[2] = y_real * freq_y_real - y_imag * freq_y_imag; - data[3] = y_real * freq_y_imag + y_imag * freq_y_real; - } - - // part 3: hadamard transform - { - // Stage 1: butterfly (data[0], data[1]) and (data[2], data[3]). - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a1; - data[1] = a0 - a1; - data[2] = a2 + a3; - data[3] = a2 - a3; - } - // Stage 2: butterfly (data[0], data[2]) and (data[1], data[3]). - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a2; - data[1] = a1 + a3; - data[2] = a0 - a2; - data[3] = a1 - a3; - } - // Stages 3..7: cross-lane butterflies. Lower-lane (mask bit clear) keeps - // the sum, upper-lane (mask bit set) keeps the difference. shfl_xor is - // unsynchronized across early-returned lanes, but invalid-plan returns - // happen above for *all* lanes of a warp (work_id is warp-uniform), so - // the warp is intact here. -#pragma unroll - for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { -#ifndef USE_ROCM - const float other = __shfl_xor_sync(kFullMask, data[i], mask, kWarpThreads); -#else - const float other = __shfl_xor(data[i], mask, kWarpThreads); -#endif - data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); - } - } - const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] *= kHadamardScale; - } - - // part 4: per-warp UE8M0 quant + store. The whole warp emits one fp8 group - // (= 128 elements) plus a single fp32 scale, matching the indexer cache - // layout (`fused_store_indexer_cache`). - { - using OutStorage = AlignedVector; - float local_max = math::abs(data[0]); -#pragma unroll - for (int i = 1; i < kVecSize; ++i) { - local_max = math::max(local_max, math::abs(data[i])); - } - const auto abs_max = warp::reduce_max(local_max); - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto inv_scale = 1.0f / scale; - const int32_t page = out_loc >> kPageBits; - const int32_t offset = out_loc & ((1 << kPageBits) - 1); - const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * 128; - const auto scale_ptr = page_ptr + (128 << kPageBits) + offset * 4; - OutStorage result; - result[0] = pack_fp8(data[0] * inv_scale, data[1] * inv_scale); - result[1] = pack_fp8(data[2] * inv_scale, data[3] * inv_scale); - PDLTriggerSecondary(); - result.store(value_ptr, lane_id); - // The single fp32 scale is identical across all lanes -- write from any lane. - if (lane_id == 0) reinterpret_cast(scale_ptr)[0] = scale; - } -} - -template -INDEXER_KERNEL void fused_norm_rope_indexer_fp4(const __grid_constant__ FusedNormRopeStoreParams params) { - using namespace device; - using enum ForwardMode; - - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kVecSize = 4; - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; - constexpr int64_t kPageBytes = 68ll << kPageBits; - static_assert(kHeadDim == kWarpThreads * kVecSize); - static_assert(kRopeDim == kWarpThreads * 2); - static_assert(kRopeSize <= kWarpThreads); - using Storage = AlignedVector; - using Float4 = AlignedVector; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kNumWarps + warp_id; - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; - - if (work_id >= params.num_tokens) return; - - const auto input = static_cast(params.input) + work_id * kHeadDim; - int32_t position; - int32_t out_loc; - if constexpr (kMode == CompressExtend) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.is_invalid()) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[plan.ragged_id]; - } else if constexpr (kMode == CompressDecode) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.seq_len % params.compress_ratio != 0) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[work_id]; - } else { - static_assert(host::dependent_false_v, "Unsupported Mode"); - } - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - PDLWaitPrimary(); - Float4 data, freq; - - { - Storage input_vec, weight_vec; - input_vec.load(input, lane_id); - weight_vec.load(params.weight, lane_id); - if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - sum_of_squares += fp32_input * fp32_input; - } - - sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); - -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - const auto fp32_weight = cast(weight_vec[i]); - data[i] = fp32_input * norm_factor * fp32_weight; - } - } - - if (is_rope_lane) { - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto y_real = data[2]; - const auto y_imag = data[3]; - const auto freq_x_real = freq[0]; - const auto freq_x_imag = freq[1]; - const auto freq_y_real = freq[2]; - const auto freq_y_imag = freq[3]; - data[0] = x_real * freq_x_real - x_imag * freq_x_imag; - data[1] = x_real * freq_x_imag + x_imag * freq_x_real; - data[2] = y_real * freq_y_real - y_imag * freq_y_imag; - data[3] = y_real * freq_y_imag + y_imag * freq_y_real; - } - - { - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a1; - data[1] = a0 - a1; - data[2] = a2 + a3; - data[3] = a2 - a3; - } - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a2; - data[1] = a1 + a3; - data[2] = a0 - a2; - data[3] = a1 - a3; - } -#pragma unroll - for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); - data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); - } - } - const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] *= kHadamardScale; - } - - { - float local_max = math::abs(data[0]); -#pragma unroll - for (int i = 1; i < kVecSize; ++i) { - local_max = math::max(local_max, math::abs(data[i])); - } - local_max = warp::reduce_max<8>(local_max); - - const auto scale_raw = fmaxf(1e-4f, local_max) / 6.0f; - const auto scale_ue8m0 = static_cast(cast_to_ue8m0(scale_raw)); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); - - const uint8_t packed0 = quant_fp4_e2m1(data[0] * inv_scale) | (quant_fp4_e2m1(data[1] * inv_scale) << 4); - const uint8_t packed1 = quant_fp4_e2m1(data[2] * inv_scale) | (quant_fp4_e2m1(data[3] * inv_scale) << 4); - const uint16_t packed = static_cast(packed0) | (static_cast(packed1) << 8); - - const int32_t page = out_loc >> kPageBits; - const int32_t offset = out_loc & ((1 << kPageBits) - 1); - const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * 64; - const auto scale_ptr = page_ptr + (64 << kPageBits) + offset * 4; - - PDLTriggerSecondary(); - reinterpret_cast(value_ptr)[lane_id] = packed; - if ((lane_id & 7) == 0) static_cast(scale_ptr)[lane_id >> 3] = scale_ue8m0; - } -} - -// ---------------------------------------------------------------------------- -// FlashMLA variant: kHeadDim = 512, 1 token per *block* (256 threads). -// Each thread loads kVecSize=2 BF16, so 256 threads cover the full 512 elems. -// Cache layout: 584 bytes/token = 448 fp8 nope + 64 (=32 bf16x2) rope + 8 scale. -// ---------------------------------------------------------------------------- -template -FLASHMLA_KERNEL void fused_norm_rope_flashmla(const __grid_constant__ FusedNormRopeStoreParams params) { - using namespace device; - using enum ForwardMode; - - constexpr int64_t kHeadDim = 512; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kVecSize = 2; - // Last warp owns the rope tail. The remaining 7 warps each emit one - // 64-element fp8 group (own UE8M0 scale). - constexpr uint32_t kRopeWarp = kNumWarps - 1; - constexpr int64_t kPageBytes = host::div_ceil(584ll << kPageBits, 576) * 576; - static_assert(kHeadDim == kBlockSize * kVecSize); - static_assert(kRopeDim == kWarpThreads * kVecSize); - static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); - using Storage = AlignedVector; - using Float2 = AlignedVector; - - const auto tx = threadIdx.x; - const auto warp_id = tx / kWarpThreads; - const auto lane_id = tx % kWarpThreads; - const auto work_id = blockIdx.x; - - if (work_id >= params.num_tokens) return; - - const auto input = static_cast(params.input) + work_id * kHeadDim; - int32_t position; - int32_t out_loc; - if constexpr (kMode == CompressExtend) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.is_invalid()) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[plan.ragged_id]; - } else if constexpr (kMode == CompressDecode) { - const auto plan = static_cast(params.handle)[work_id]; - if (plan.seq_len % params.compress_ratio != 0) return; - position = plan.seq_len - params.compress_ratio; - out_loc = params.out_loc[work_id]; - } else { - static_assert(host::dependent_false_v, "Unsupported Mode"); - } - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - PDLWaitPrimary(); - Float2 data, freq; - - // part 1: norm. Each thread owns one 2-elem pack (`tx`-th pack of input). - // Sum of squares is reduced across the whole block via per-warp partials. - { - __shared__ float partial_sums[kNumWarps]; - - Storage input_vec, weight_vec; - input_vec.load(input, tx); - weight_vec.load(params.weight, tx); - if (warp_id == kRopeWarp) freq.load(freqs_cis, lane_id); - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - sum_of_squares += fp32_input * fp32_input; - } - - const auto warp_sum = warp::reduce_sum(sum_of_squares); - if (lane_id == 0) partial_sums[warp_id] = warp_sum; - __syncthreads(); - // Replicate the per-warp partial sums to a full warp and reduce. Every - // lane-group of `kNumWarps` lanes ends up with the global sum. - sum_of_squares = warp::reduce_sum(partial_sums[lane_id % kNumWarps]); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); - -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto fp32_input = cast(input_vec[i]); - const auto fp32_weight = cast(weight_vec[i]); - data[i] = fp32_input * norm_factor * fp32_weight; - } - } - - const int32_t page = out_loc >> kPageBits; - const int32_t offset = out_loc & ((1 << kPageBits) - 1); - const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * 576; - - PDLTriggerSecondary(); - - // part 2: rope on the rope warp (BF16 store), or per-warp FP8 quant + store. - if (warp_id == kRopeWarp) { - // Each rope-warp lane owns exactly one (real, imag) pair within the rope - // tail. Apply rotation, downcast to BF16, write to the slot's rope region. - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto freq_real = freq[0]; - const auto freq_imag = freq[1]; - data[0] = x_real * freq_real - x_imag * freq_imag; - data[1] = x_real * freq_imag + x_imag * freq_real; - const auto result = cast(fp32x2_t{data[0], data[1]}); - const auto rope_ptr = value_ptr + 448; - reinterpret_cast(rope_ptr)[lane_id] = result; - } else { - // Non-rope warp: per-warp UE8M0 group (64 elems -> 64 fp8 + 1 scale byte). - // BF16 round-trip to match the precision of the non-fused path - // (which goes through quant_to_nope_fp8_rope_bf16_pack_triton with bf16 input). - const auto x = cast(cast(data[0])); - const auto y = cast(cast(data[1])); - const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); - const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); - const auto result = pack_fp8(x * inv_scale, y * inv_scale); - const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; - reinterpret_cast(value_ptr)[tx] = result; - // All lanes in this warp produce the same scale byte; let lane 0 publish. - if (lane_id == 0) static_cast(scale_ptr)[warp_id] = scale_ue8m0; - } -} - -template -struct FusedNormRopeKernel { - static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); - static constexpr bool kIsIndexer = (kHeadDim == 128); - static constexpr int64_t kIndexerBytes = 132 * kPageSize; - static constexpr int64_t kFlashMLABytes = host::div_ceil(584 * kPageSize, 576) * 576; - static constexpr int64_t kPageBytes = kIsIndexer ? kIndexerBytes : kFlashMLABytes; - - /// TODO: Let's fix the config for now. - static_assert(kRopeDim == 64 && (kHeadDim == 128 || kHeadDim == 512)); - static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); - - template - static constexpr auto select_kernel() { - if constexpr (kIsIndexer) { - return fused_norm_rope_indexer; - } else { - return fused_norm_rope_flashmla; - } - } - - template - static constexpr auto select_fp4_kernel() { - static_assert(kIsIndexer, "FP4 fused store is only defined for the indexer"); - return fused_norm_rope_indexer_fp4; - } - - static void forward( - const tvm::ffi::TensorView input, - const tvm::ffi::TensorView plan, - const tvm::ffi::TensorView weight, - const float eps, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView out_loc, - const tvm::ffi::TensorView kvcache, - const bool is_decode, - const uint32_t compress_ratio) { - using namespace host; - using enum ForwardMode; - - const auto mode = static_cast(is_decode); - - auto N = SymbolicSize{"num_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({N, kHeadDim}) // input - .with_dtype() - .with_device(device_) - .verify(input); - TensorMatcher({kHeadDim}) // weight - .with_dtype() - .with_device(device_) - .verify(weight); - TensorMatcher({-1, kRopeDim}) // freqs_cis - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - TensorMatcher({-1}) // out_loc - .with_dtype() - .with_device(device_) - .verify(out_loc); - TensorMatcher({-1, -1}) // cache - .with_strides({kPageBytes, 1}) - .with_dtype() - .with_device(device_) - .verify(kvcache); - - switch (mode) { - case CompressExtend: - compress::verify_plan_c(plan, N, device_); - RuntimeCheck(out_loc.size(0) >= N.unwrap()); - break; - case CompressDecode: - compress::verify_plan_d(plan, N, device_); - RuntimeCheck(out_loc.size(0) == N.unwrap()); - break; - } - - const auto num_tokens = static_cast(N.unwrap()); - if (num_tokens == 0) return; - const auto params = FusedNormRopeStoreParams{ - .input = input.data_ptr(), - .handle = plan.data_ptr(), - .weight = weight.data_ptr(), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .out_loc = static_cast(out_loc.data_ptr()), - .kvcache = static_cast(kvcache.data_ptr()), - .eps = eps, - .compress_ratio = compress_ratio, - .num_tokens = num_tokens, - }; - // Indexer packs `kNumWarps` tokens per block (warp-major); FlashMLA uses - // a whole block per token (cta-major sum-reduce over head_dim=512). - const uint32_t num_blocks = kIsIndexer ? div_ceil(num_tokens, kNumWarps) : num_tokens; - const auto device = device_.unwrap(); - const auto kernel = mode == CompressExtend ? select_kernel() : select_kernel(); - LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params); - } - - static void forward_fp4( - const tvm::ffi::TensorView input, - const tvm::ffi::TensorView plan, - const tvm::ffi::TensorView weight, - const float eps, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView out_loc, - const tvm::ffi::TensorView kvcache, - const bool is_decode, - const uint32_t compress_ratio) { - using namespace host; - using enum ForwardMode; - - static_assert(kIsIndexer, "FP4 fused store is only defined for the indexer"); - constexpr int64_t kFp4PageBytes = 68 * kPageSize; - const auto mode = static_cast(is_decode); - - auto N = SymbolicSize{"num_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({N, kHeadDim}).with_dtype().with_device(device_).verify(input); - TensorMatcher({kHeadDim}).with_dtype().with_device(device_).verify(weight); - TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); - TensorMatcher({-1}).with_dtype().with_device(device_).verify(out_loc); - TensorMatcher({-1, -1}).with_strides({kFp4PageBytes, 1}).with_dtype().with_device(device_).verify(kvcache); - - switch (mode) { - case CompressExtend: - compress::verify_plan_c(plan, N, device_); - RuntimeCheck(out_loc.size(0) >= N.unwrap()); - break; - case CompressDecode: - compress::verify_plan_d(plan, N, device_); - RuntimeCheck(out_loc.size(0) == N.unwrap()); - break; - } - - const auto num_tokens = static_cast(N.unwrap()); - if (num_tokens == 0) return; - const auto params = FusedNormRopeStoreParams{ - .input = input.data_ptr(), - .handle = plan.data_ptr(), - .weight = weight.data_ptr(), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .out_loc = static_cast(out_loc.data_ptr()), - .kvcache = static_cast(kvcache.data_ptr()), - .eps = eps, - .compress_ratio = compress_ratio, - .num_tokens = num_tokens, - }; - const uint32_t num_blocks = div_ceil(num_tokens, kNumWarps); - const auto device = device_.unwrap(); - const auto kernel = - mode == CompressExtend ? select_fp4_kernel() : select_fp4_kernel(); - LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh deleted file mode 100644 index 90dec3c117..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hash_topk.cuh +++ /dev/null @@ -1,214 +0,0 @@ -#include -#include - -#include -#include -#include - -#include - -#include -#include - -namespace { - -[[maybe_unused]] -SGL_DEVICE float act_sqrt_softplus(float x) { - const float softplus = fmaxf(x, 0.0f) + log1pf(expf(-fabsf(x))); - return sqrtf(softplus); -} - -struct MoEHashTopKParams { - const float* __restrict__ router_logits; - const int64_t* __restrict__ input_id; - const int32_t* __restrict__ tid2eid; - int32_t* __restrict__ topk_ids; - float* __restrict__ topk_weights; - uint32_t num_tokens; - uint32_t topk; - uint32_t num_routed_experts; - uint32_t num_shared_experts; - float routed_scaling_factor; -}; - -template -__global__ void moe_hash_topk_fused(const MoEHashTopKParams __grid_constant__ params) { - using namespace device; - const auto& [ - router_logits, input_id, tid2eid, topk_ids, topk_weights, // pointers - num_tokens, topk, num_routed_experts, num_shared_experts, routed_scaling_factor] = - params; - - const uint32_t topk_fused = topk + num_shared_experts; - const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t warp_id = tid / kWarpThreads; - const uint32_t lane_id = tid % kWarpThreads; - if (warp_id >= num_tokens) return; - // we can safely prefetch the token id - const auto token_id = input_id[warp_id]; - - PDLWaitPrimary(); - - float routed_weight = 0.0f; - int32_t expert_id = 0; - if (lane_id < topk) { - expert_id = tid2eid[token_id * topk + lane_id]; - routed_weight = Fn(router_logits[warp_id * num_routed_experts + expert_id]); - } - - const auto routed_sum = device::warp::reduce_sum(routed_weight); - if (lane_id < topk_fused) { - const bool is_shared = lane_id >= topk; - const auto output_offset = warp_id * topk_fused + lane_id; - topk_ids[output_offset] = is_shared ? num_routed_experts + lane_id - topk : expert_id; - topk_weights[output_offset] = is_shared ? 1.0f / routed_scaling_factor : routed_weight / routed_sum; - } - - PDLTriggerSecondary(); -} - -struct TopKParams { - int32_t* __restrict__ topk_ids; - // Exactly one is active: ntn_ptr == nullptr means use ntn_value. - const int32_t* __restrict__ ntn_ptr; - int32_t ntn_value; - int64_t stride; - uint32_t topk; - uint32_t num_tokens; -}; - -__global__ void mask_topk_ids_padded_region(const TopKParams __grid_constant__ params) { - const uint32_t tid = blockIdx.x * blockDim.x + threadIdx.x; - const uint32_t warp_id = tid / device::kWarpThreads; - const uint32_t lane_id = tid % device::kWarpThreads; - if (warp_id >= params.num_tokens || lane_id >= params.topk) return; - device::PDLWaitPrimary(); - const uint32_t num = (params.ntn_ptr != nullptr) // - ? static_cast(params.ntn_ptr[0]) - : static_cast(params.ntn_value); - if (warp_id >= num) params.topk_ids[warp_id * params.stride + lane_id] = -1; - device::PDLTriggerSecondary(); -} - -template -struct HashTopKKernel { - static constexpr auto kernel = moe_hash_topk_fused; - - static void - run(const tvm::ffi::TensorView router_logits, - const tvm::ffi::TensorView input_id, - const tvm::ffi::TensorView tid2eid, - const tvm::ffi::TensorView topk_weights, - const tvm::ffi::TensorView topk_ids, - float routed_scaling_factor) { - using namespace host; - - auto N = SymbolicSize{"num_tokens"}; - auto E = SymbolicSize{"num_routed_experts"}; - auto K = SymbolicSize{"topk_fused"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({N, E}) // - .with_dtype() - .with_device(device) - .verify(router_logits); - TensorMatcher({N}) // - .with_dtype() - .with_device(device) - .verify(input_id); - TensorMatcher({-1, -1}) // - .with_dtype() - .with_device(device) - .verify(tid2eid); - TensorMatcher({N, K}) // - .with_dtype() - .with_device(device) - .verify(topk_weights); - TensorMatcher({N, K}) // - .with_dtype() - .with_device(device) - .verify(topk_ids); - - const auto num_tokens = static_cast(N.unwrap()); - const auto topk_fused = static_cast(K.unwrap()); - const auto topk = static_cast(tid2eid.size(1)); - const auto shared_experts = topk_fused - topk; - RuntimeCheck(topk <= topk_fused, "HashTopKKernel requires topk <= topk_fused"); - RuntimeCheck(topk_fused <= device::kWarpThreads, "HashTopKKernel requires topk_fused <= warp size"); - - const auto params = MoEHashTopKParams{ - .router_logits = static_cast(router_logits.data_ptr()), - .input_id = static_cast(input_id.data_ptr()), - .tid2eid = static_cast(tid2eid.data_ptr()), - .topk_ids = static_cast(topk_ids.data_ptr()), - .topk_weights = static_cast(topk_weights.data_ptr()), - .num_tokens = num_tokens, - .topk = topk, - .num_routed_experts = static_cast(E.unwrap()), - .num_shared_experts = shared_experts, - .routed_scaling_factor = routed_scaling_factor, - }; - const auto kBlockSize = 128u; - const auto kNumWarps = kBlockSize / device::kWarpThreads; - const auto num_blocks = div_ceil(num_tokens, kNumWarps); - LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -// TODO this may not be related to *hash* topk, thus may move -struct MaskKernel { - static constexpr auto kernel = mask_topk_ids_padded_region; - - static void run(tvm::ffi::TensorView topk_ids, tvm::ffi::TensorView num_token_non_padded) { - using namespace host; - - auto N = SymbolicSize{"num_tokens"}; - auto K = SymbolicSize{"topk"}; - auto D = SymbolicSize{"stride"}; - auto device = SymbolicDevice{}; - device.set_options(); - TensorMatcher({N, K}) // - .with_strides({D, 1}) - .with_dtype() - .with_device(device) - .verify(topk_ids); - RuntimeCheck(num_token_non_padded.numel() == 1, "num_token_non_padded should be a scalar"); - RuntimeCheck(K.unwrap() <= device::kWarpThreads, "MaskKernel requires topk <= warp size"); - const int32_t* ntn_ptr = nullptr; - int32_t ntn_value = 0; - const auto ntn_dev = num_token_non_padded.device().device_type; - if (ntn_dev == kDLCUDA) { - RuntimeCheck(is_type(num_token_non_padded.dtype()), "num_token_non_padded on CUDA must be int32"); - ntn_ptr = static_cast(num_token_non_padded.data_ptr()); - } else if (ntn_dev == kDLCPU) { - if (is_type(num_token_non_padded.dtype())) { - ntn_value = *static_cast(num_token_non_padded.data_ptr()); - } else if (is_type(num_token_non_padded.dtype())) { - ntn_value = static_cast(*static_cast(num_token_non_padded.data_ptr())); - } else { - RuntimeCheck(false, "num_token_non_padded on CPU must be int32 or int64"); - } - } else { - RuntimeCheck(false, "num_token_non_padded must be on CPU or CUDA"); - } - - const auto num_tokens = static_cast(N.unwrap()); - const auto params = TopKParams{ - .topk_ids = static_cast(topk_ids.data_ptr()), - .ntn_ptr = ntn_ptr, - .ntn_value = ntn_value, - .stride = static_cast(D.unwrap()), - .topk = static_cast(K.unwrap()), - .num_tokens = num_tokens, - }; - const auto kBlockSize = 128u; - const auto kNumWarps = kBlockSize / device::kWarpThreads; - const auto num_blocks = div_ceil(num_tokens, kNumWarps); - LaunchKernel(num_blocks, kBlockSize, device.unwrap()) // - .enable_pdl(true)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh deleted file mode 100644 index aefec24372..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/hisparse_transfer.cuh +++ /dev/null @@ -1,82 +0,0 @@ -#include -#include - -#include - -#include - -#include -#include - -#include - -namespace { - -/// NOTE: for offload to cpu kernel, we use persistent kernel -inline constexpr uint32_t kBlockSize = 1024; -inline constexpr uint32_t kBlockQuota = 4; - -#define OFFLOAD_KERNEL __global__ __launch_bounds__(kBlockSize, 1) - -struct OffloadParams { - void** gpu_caches; - void** cpu_caches; - const int64_t* gpu_indices; - const int64_t* cpu_indices; - uint32_t num_items; - uint32_t num_layers; -}; - -OFFLOAD_KERNEL void offload_to_cpu(const __grid_constant__ OffloadParams params) { - using namespace device::hisparse; - const auto [gpu_caches, cpu_caches, gpu_indices, cpu_indices, num_items, num_layers] = params; - const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; - constexpr auto kNumWarps = (kBlockSize / 32) * kBlockQuota; - for (auto i = global_tid / 32; i < num_items; i += kNumWarps) { - const int32_t gpu_index = gpu_indices[i]; - const int32_t cpu_index = cpu_indices[i]; - for (auto j = 0u; j < num_layers; ++j) { - const auto gpu_cache = gpu_caches[j]; - const auto cpu_cache = cpu_caches[j]; - transfer_item( - /*dst_cache=*/cpu_cache, - /*src_cache=*/gpu_cache, - /*dst_index=*/cpu_index, - /*src_index=*/gpu_index); - } - } -} - -[[maybe_unused]] -void hisparse_transfer( - tvm::ffi::TensorView gpu_ptrs, - tvm::ffi::TensorView cpu_ptrs, - tvm::ffi::TensorView gpu_indices, - tvm::ffi::TensorView cpu_indices) { - using namespace host; - auto N = SymbolicSize{"num_items"}; - auto L = SymbolicSize{"num_layers"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - TensorMatcher({L}) // 1D cache pointers - .with_dtype() - .with_device(device_) - .verify(gpu_ptrs) - .verify(cpu_ptrs); - TensorMatcher({N}) // 1D indices - .with_dtype() - .with_device(device_) - .verify(gpu_indices) - .verify(cpu_indices); - const auto params = OffloadParams{ - .gpu_caches = static_cast(gpu_ptrs.data_ptr()), - .cpu_caches = static_cast(cpu_ptrs.data_ptr()), - .gpu_indices = static_cast(gpu_indices.data_ptr()), - .cpu_indices = static_cast(cpu_indices.data_ptr()), - .num_items = static_cast(N.unwrap()), - .num_layers = static_cast(L.unwrap()), - }; - LaunchKernel(kBlockQuota, kBlockSize, device_.unwrap())(offload_to_cpu, params); -} - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh deleted file mode 100644 index 8fc8d0821d..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/main_norm_rope.cuh +++ /dev/null @@ -1,845 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include - -#include -#include - -namespace { - -using deepseek_v4::fp8::cast_to_ue8m0; -using deepseek_v4::fp8::inv_scale_ue8m0; -using deepseek_v4::fp8::pack_fp8; - -SGL_DEVICE uint8_t quant_fp4_e2m1(float x) { - const float ax = fminf(fabsf(x), 6.0f); - uint8_t idx = 0; - idx += ax > 0.25f; - idx += ax > 0.75f; - idx += ax > 1.25f; - idx += ax > 1.75f; - idx += ax > 2.5f; - idx += ax > 3.5f; - idx += ax > 5.0f; - if (x < 0.0f && idx != 0) idx |= 0x8; - return idx; -} - -// 4 warps per block: warp-per-(token, head) work-item dispatch (Q kernel). -constexpr uint32_t kFusedQBlockSize = 128; -constexpr uint32_t kFusedQNumWarps = kFusedQBlockSize / device::kWarpThreads; - -// 8 warps per block: block-per-token work-item dispatch (K kernel). -constexpr uint32_t kFusedKBlockSize = 256; -constexpr uint32_t kFusedKNumWarps = kFusedKBlockSize / device::kWarpThreads; - -#define Q_KERNEL __global__ __launch_bounds__(kFusedQBlockSize, 16) -#define K_KERNEL __global__ __launch_bounds__(kFusedKBlockSize, 8) - -// ============================================================================ -// Q kernel: warp-per-(token, head) rmsnorm-self + RoPE + write to q_out. -// ============================================================================ - -struct FusedQNormRopeParams { - const void* __restrict__ q_input; // (B, num_q_heads, kHeadDim) DType - void* __restrict__ q_output; // (B, num_q_heads, kHeadDim) DType - const float* __restrict__ freqs_cis; // (max_pos, kRopeDim) fp32 (re/im interleaved) - const void* __restrict__ positions; // (B,) PosT - int64_t q_input_stride_batch; - int64_t q_output_stride_batch; - uint32_t batch_size; - uint32_t num_q_heads; - float eps; -}; - -template -Q_KERNEL void fused_q_norm_rope(const __grid_constant__ FusedQNormRopeParams params) { - using namespace device; - - constexpr int64_t kMaxVecSize = 16 / sizeof(DType); - constexpr int64_t kVecSize = std::min(kMaxVecSize, kHeadDim / kWarpThreads); - constexpr int64_t kLocalSize = kHeadDim / (kWarpThreads * kVecSize); - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; - static_assert(kHeadDim % (kWarpThreads * kVecSize) == 0); - static_assert(kLocalSize * kVecSize * kWarpThreads == kHeadDim); - static_assert(kRopeDim % kVecSize == 0); - static_assert(kRopeSize <= kWarpThreads); - static_assert(kRopeDim == kWarpThreads * 2, "1 (real, imag) pair per lane"); - - using Storage = AlignedVector; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; - - const uint32_t total_works = params.batch_size * params.num_q_heads; - if (work_id >= total_works) return; - - const uint32_t batch_id = work_id / params.num_q_heads; - const uint32_t head_id = work_id % params.num_q_heads; - const auto input_ptr = - static_cast(params.q_input) + batch_id * params.q_input_stride_batch + head_id * kHeadDim; - const auto output_ptr = - static_cast(params.q_output) + batch_id * params.q_output_stride_batch + head_id * kHeadDim; - const auto position = static_cast(static_cast(params.positions)[batch_id]); - - __shared__ Storage s_rope[kFusedQNumWarps][kRopeSize]; - - // Prefetch this lane's freq pair before the PDL gate so the wait happens - // outside the dependency chain on `position`. - const auto mem_freq = tile::Memory{lane_id, kWarpThreads}; - - PDLWaitPrimary(); - - // part 1: rmsnorm-self (no weight). - const auto gmem = tile::Memory{lane_id, kWarpThreads}; - Storage input_vec[kLocalSize]; -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { - input_vec[i] = gmem.load(input_ptr, i); - } - - const auto freq = mem_freq.load(params.freqs_cis + position * kRopeDim); - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { -#pragma unroll - for (int j = 0; j < kVecSize; ++j) { - const auto x = cast(input_vec[i][j]); - sum_of_squares += x * x; - } - } - sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); - -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { -#pragma unroll - for (int j = 0; j < kVecSize; ++j) { - const auto x = cast(input_vec[i][j]); - input_vec[i][j] = cast(x * norm_factor); - } - } - - // Stash the rope tail (last kRopeSize lanes' last tile) into shared memory; - // write nope tiles to gmem directly. - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; -#pragma unroll - for (int i = 0; i < kLocalSize; ++i) { - if (i == kLocalSize - 1 && is_rope_lane) { - const auto rope_id = lane_id - (kWarpThreads - kRopeSize); - s_rope[warp_id][rope_id] = input_vec[i]; - } else { - gmem.store(output_ptr, input_vec[i], i); - } - } - __syncwarp(); - - PDLTriggerSecondary(); - - // part 2: RoPE on all 32 lanes -- one (real, imag) bf16x2 pair per lane. - using DType2 = packed_t; - const auto mem_elem = tile::Memory{lane_id, kWarpThreads}; - const auto elem = mem_elem.load(s_rope[warp_id]); - const auto [x_real, x_imag] = cast(elem); - const auto [freq_real, freq_imag] = freq; - const fp32x2_t rotated = { - x_real * freq_real - x_imag * freq_imag, - x_real * freq_imag + x_imag * freq_real, - }; - mem_elem.store(output_ptr + (kHeadDim - kRopeDim), cast(rotated)); -} - -template -struct FusedQNormRopeKernel { - template - static constexpr auto kernel = fused_q_norm_rope; - - static void forward( - const tvm::ffi::TensorView q_input, - const tvm::ffi::TensorView q_output, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView positions, - float eps) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto H = SymbolicSize{"num_q_heads"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, H, kHeadDim}) // - .with_strides({-1, kHeadDim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_input); - TensorMatcher({B, H, kHeadDim}) // - .with_strides({-1, kHeadDim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_output); - TensorMatcher({-1, kRopeDim}) // - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - auto pos_dtype = SymbolicDType{}; - TensorMatcher({B}) // - .with_dtype(pos_dtype) - .with_device(device_) - .verify(positions); - - const auto batch_size = static_cast(B.unwrap()); - const auto num_q_heads = static_cast(H.unwrap()); - if (batch_size == 0) return; - - const auto params = FusedQNormRopeParams{ - .q_input = q_input.data_ptr(), - .q_output = q_output.data_ptr(), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .positions = positions.data_ptr(), - .q_input_stride_batch = q_input.stride(0), - .q_output_stride_batch = q_output.stride(0), - .batch_size = batch_size, - .num_q_heads = num_q_heads, - .eps = eps, - }; - const auto total_works = batch_size * num_q_heads; - const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); - const auto k_int32 = kernel; - const auto k_int64 = kernel; - const auto k = pos_dtype.is_type() ? k_int32 : k_int64; - LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(k, params); - } -}; - -// ============================================================================ -// K kernel: block-per-token rmsnorm (with kv_weight) + RoPE + FlashMLA store. -// ============================================================================ - -struct FusedKNormRopeFlashMLAParams { - const void* __restrict__ kv; // (B, kHeadDim) DType - const void* __restrict__ kv_weight; // (kHeadDim,) DType - const float* __restrict__ freqs_cis; // (max_pos, kRopeDim) fp32 - const void* __restrict__ positions; // (B,) PosT - const int32_t* __restrict__ out_loc; // (B,) int32 -> cache slot id - uint8_t* __restrict__ kvcache; // (npages, kPageBytes) uint8 - // Row stride for `kv` in elements. Required because the upstream caller often - // passes `qkv_a[..., q_lora_rank:]`, a non-contiguous slice whose stride[0] - // equals `q_lora_rank + kHeadDim` rather than `kHeadDim`. - int64_t kv_stride_batch; - uint32_t batch_size; - float eps; -}; - -template -K_KERNEL void fused_k_norm_rope_flashmla(const __grid_constant__ FusedKNormRopeFlashMLAParams params) { - using namespace device; - - constexpr int64_t kVecSize = 2; - constexpr uint32_t kRopeWarp = kFusedKNumWarps - 1; - constexpr int64_t kPageBytes = host::div_ceil(584ll << kPageBits, 576) * 576; - static_assert(kHeadDim == kFusedKBlockSize * kVecSize); - static_assert(kRopeDim == kWarpThreads * kVecSize); - static_assert(kHeadDim - kRopeDim == kRopeWarp * kWarpThreads * kVecSize); - using Storage = AlignedVector; - using Float2 = AlignedVector; - - const auto tx = threadIdx.x; - const auto warp_id = tx / kWarpThreads; - const auto lane_id = tx % kWarpThreads; - const auto work_id = blockIdx.x; - if (work_id >= params.batch_size) return; - - const auto input_ptr = static_cast(params.kv) + work_id * params.kv_stride_batch; - const auto position = static_cast(static_cast(params.positions)[work_id]); - const auto out_loc = params.out_loc[work_id]; - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - PDLWaitPrimary(); - Float2 data, freq; - - // part 1: norm. Each thread owns one 2-elem pack (the `tx`-th). - // Sum-of-squares is reduced block-wide via per-warp partials. - { - __shared__ float partial_sums[kFusedKNumWarps]; - - Storage input_vec, weight_vec; - input_vec.load(input_ptr, tx); - weight_vec.load(params.kv_weight, tx); - if (warp_id == kRopeWarp) freq.load(freqs_cis, lane_id); - - float sum_of_squares = 0.0f; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto x = cast(input_vec[i]); - sum_of_squares += x * x; - } - const auto warp_sum = warp::reduce_sum(sum_of_squares); - if (lane_id == 0) partial_sums[warp_id] = warp_sum; - __syncthreads(); - // Replicate the per-warp partial sums onto all lanes of one warp and - // reduce. Every group of `kBlockItemNumWarps` lanes ends up with the - // global sum. - sum_of_squares = warp::reduce_sum(partial_sums[lane_id % kFusedKNumWarps]); - const auto norm_factor = math::rsqrt(sum_of_squares / kHeadDim + params.eps); - -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const auto x = cast(input_vec[i]); - const auto w = cast(weight_vec[i]); - data[i] = x * norm_factor * w; - } - } - - const int32_t page = out_loc >> kPageBits; - const int32_t offset = out_loc & ((1 << kPageBits) - 1); - const auto page_ptr = params.kvcache + page * kPageBytes; - const auto value_ptr = page_ptr + offset * 576; - - PDLTriggerSecondary(); - - // part 2: rope on warp 7 (BF16 store), per-warp UE8M0 quant + store on warps 0..6. - if (warp_id == kRopeWarp) { - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto freq_real = freq[0]; - const auto freq_imag = freq[1]; - data[0] = x_real * freq_real - x_imag * freq_imag; - data[1] = x_real * freq_imag + x_imag * freq_real; - const auto result = cast(fp32x2_t{data[0], data[1]}); - const auto rope_ptr = value_ptr + 448; - reinterpret_cast(rope_ptr)[lane_id] = result; - } else { - const auto x = data[0]; - const auto y = data[1]; - const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); - const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); - const auto result = pack_fp8(x * inv_scale, y * inv_scale); - const auto scale_ptr = page_ptr + (576 << kPageBits) + offset * 8; - reinterpret_cast(value_ptr)[tx] = result; - if (lane_id == 0) static_cast(scale_ptr)[warp_id] = scale_ue8m0; - } -} - -template -struct FusedKNormRopeFlashMLAKernel { - static constexpr int32_t kLogPageSize = std::countr_zero(kPageSize); - static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; - static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); - static_assert(1 << kLogPageSize == kPageSize); - static_assert(kHeadDim == 512 && kRopeDim == 64, "FlashMLA layout requires (512, 64)"); - - template - static constexpr auto kernel = fused_k_norm_rope_flashmla; - - static void forward( - const tvm::ffi::TensorView kv, - const tvm::ffi::TensorView kv_weight, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView positions, - const tvm::ffi::TensorView out_loc, - const tvm::ffi::TensorView kvcache, - float eps) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, kHeadDim}) // - .with_strides({-1, 1}) - .with_dtype() - .with_device(device_) - .verify(kv); - TensorMatcher({kHeadDim}) // - .with_dtype() - .with_device(device_) - .verify(kv_weight); - TensorMatcher({-1, kRopeDim}) // - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - auto pos_dtype = SymbolicDType{}; - TensorMatcher({B}) // - .with_dtype(pos_dtype) - .with_device(device_) - .verify(positions); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(out_loc); - TensorMatcher({-1, -1}) // - .with_strides({kPageBytes, 1}) - .with_dtype() - .with_device(device_) - .verify(kvcache); - - const auto batch_size = static_cast(B.unwrap()); - if (batch_size == 0) return; - - const auto params = FusedKNormRopeFlashMLAParams{ - .kv = kv.data_ptr(), - .kv_weight = kv_weight.data_ptr(), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .positions = positions.data_ptr(), - .out_loc = static_cast(out_loc.data_ptr()), - .kvcache = static_cast(kvcache.data_ptr()), - .kv_stride_batch = kv.stride(0), - .batch_size = batch_size, - .eps = eps, - }; - const auto k_int32 = kernel; - const auto k_int64 = kernel; - const auto k = pos_dtype.is_type() ? k_int32 : k_int64; - LaunchKernel(batch_size, kFusedKBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(k, params); - } -}; - -// ============================================================================ -// Indexer Q kernel: warp-per-(token, head) RoPE + Hadamard + fp8 act-quant. -// ============================================================================ - -struct FusedQIndexerRopeHadamardQuantParams { - const void* __restrict__ q_input; // (B, num_heads, 128) DType - void* __restrict__ q_fp8; // (B, num_heads, 128) fp8_e4m3 - // weights_out[b, h] = weight[b, h] * weight_scale * q_scale[b, h]. - // q_scale is computed internally and not exposed -- the only consumer of - // it is `weights_out`. - const void* __restrict__ weight; // (B, num_heads) DType - float* __restrict__ weights_out; // (B, num_heads) fp32 (== (B, H, 1) flat) - float weight_scale; // scalar c4_indexer.weight_scale - const float* __restrict__ freqs_cis; // (max_pos, 64) fp32 - const void* __restrict__ positions; // (B,) PosT - uint32_t batch_size; - uint32_t num_heads; -}; - -template -Q_KERNEL void fused_q_indexer_rope_hadamard_quant(const __grid_constant__ FusedQIndexerRopeHadamardQuantParams params) { - using namespace device; - - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kVecSize = 4; - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; // = 16 - static_assert(kHeadDim == kWarpThreads * kVecSize); - static_assert(kRopeDim == kWarpThreads * 2); - static_assert(kRopeSize <= kWarpThreads); - - using Storage = AlignedVector; - using Float4 = AlignedVector; - using OutStorage = AlignedVector; // 4 fp8 / lane - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; - // Last `kRopeSize` lanes own the rope tail; their 4-elem packs cover the - // trailing kRopeDim elements. - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; - - const uint32_t total_works = params.batch_size * params.num_heads; - if (work_id >= total_works) return; - - const uint32_t batch_id = work_id / params.num_heads; - const auto input_ptr = static_cast(params.q_input) + work_id * kHeadDim; - const auto position = static_cast(static_cast(params.positions)[batch_id]); - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - // Lane 0 prefetches the weight scalar for this (token, head) work item. - // Weight is (B, num_heads) DType; we need one scalar per warp -- offload - // the load to lane 0 only. The multiply + store happens once the q_scale - // is known (part 4). - - PDLWaitPrimary(); - Float4 data, freq; - const auto weight_val = cast(static_cast(params.weight)[work_id]); - - // part 1: load (no norm). Each lane owns a 4-elem pack. - { - Storage input_vec; - input_vec.load(input_ptr, lane_id); - if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - data[i] = cast(input_vec[i]); - } - } - - // part 2: rope on rope lanes only (4 elems / lane = 2 (real, imag) pairs). - if (is_rope_lane) { - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto y_real = data[2]; - const auto y_imag = data[3]; - const auto fxr = freq[0]; - const auto fxi = freq[1]; - const auto fyr = freq[2]; - const auto fyi = freq[3]; - data[0] = x_real * fxr - x_imag * fxi; - data[1] = x_real * fxi + x_imag * fxr; - data[2] = y_real * fyr - y_imag * fyi; - data[3] = y_real * fyi + y_imag * fyr; - } - - PDLTriggerSecondary(); - - // part 3: 128-point Hadamard (2 local stages + 5 cross-lane shfl_xor stages). - // Same recipe as `fused_norm_rope_indexer`; see comments there for the - // butterfly invariants and the early-return safety argument. - { - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a1; - data[1] = a0 - a1; - data[2] = a2 + a3; - data[3] = a2 - a3; - } - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a2; - data[1] = a1 + a3; - data[2] = a0 - a2; - data[3] = a1 - a3; - } -#pragma unroll - for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); - data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); - } - } - const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] *= kHadamardScale; - } - - { - float local_max = math::abs(data[0]); -#pragma unroll - for (int i = 1; i < kVecSize; ++i) { - local_max = math::max(local_max, math::abs(data[i])); - } - const auto abs_max = warp::reduce_max(local_max); - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto inv_scale = 1.0f / scale; - OutStorage result; - result[0] = pack_fp8(data[0] * inv_scale, data[1] * inv_scale); - result[1] = pack_fp8(data[2] * inv_scale, data[3] * inv_scale); - - // q_fp8 row pointer: 128 fp8 / row = 32 OutStorage / row, one per lane. - auto out_row = static_cast(params.q_fp8) + work_id * kHeadDim; - result.store(out_row, lane_id); - params.weights_out[work_id] = weight_val * params.weight_scale * scale; - } -} - -template -struct FusedQIndexerRopeHadamardQuantKernel { - template - static constexpr auto kernel = fused_q_indexer_rope_hadamard_quant; - - static void forward( - const tvm::ffi::TensorView q_input, - const tvm::ffi::TensorView q_fp8, - const tvm::ffi::TensorView weight, - const tvm::ffi::TensorView weights_out, - double weight_scale, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView positions) { - using namespace host; - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - - auto B = SymbolicSize{"batch_size"}; - auto H = SymbolicSize{"num_heads"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - // Caller path is `wq_b(q_lora).view(-1, H, D)` -> contiguous; the kernel - // assumes a flat `(B*H, kHeadDim)` layout for both q_input and q_fp8. - // Pin the head/innermost strides; assert the batch stride below. - TensorMatcher({B, H, kHeadDim}) // - .with_strides({-1, kHeadDim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_input); - TensorMatcher({B, H, kHeadDim}) // - .with_strides({-1, kHeadDim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_fp8); - TensorMatcher({B, H}) // - .with_dtype() - .with_device(device_) - .verify(weight); - TensorMatcher({B, H, 1}) // - .with_dtype() - .with_device(device_) - .verify(weights_out); - TensorMatcher({-1, kRopeDim}) // - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - auto pos_dtype = SymbolicDType{}; - TensorMatcher({B}) // - .with_dtype(pos_dtype) - .with_device(device_) - .verify(positions); - - const auto batch_size = static_cast(B.unwrap()); - const auto num_heads = static_cast(H.unwrap()); - if (batch_size == 0) return; - - // The kernel computes row pointers as `base + work_id * kHeadDim`, so - // both inputs must be contiguous in (batch, head, elem) order. - const int64_t expected_batch_stride = static_cast(num_heads) * kHeadDim; - RuntimeCheck( - q_input.stride(0) == expected_batch_stride, - "q_input must be contiguous (B, H, kHeadDim); got stride[0]=", - q_input.stride(0)); - RuntimeCheck( - q_fp8.stride(0) == expected_batch_stride, - "q_fp8 must be contiguous (B, H, kHeadDim); got stride[0]=", - q_fp8.stride(0)); - - const auto params = FusedQIndexerRopeHadamardQuantParams{ - .q_input = q_input.data_ptr(), - .q_fp8 = q_fp8.data_ptr(), - .weight = weight.data_ptr(), - .weights_out = static_cast(weights_out.data_ptr()), - .weight_scale = static_cast(weight_scale), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .positions = positions.data_ptr(), - .batch_size = batch_size, - .num_heads = num_heads, - }; - const auto total_works = batch_size * num_heads; - const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); - const auto k_int32 = kernel; - const auto k_int64 = kernel; - const auto k = pos_dtype.is_type() ? k_int32 : k_int64; - LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(k, params); - } -}; - -struct FusedQIndexerRopeHadamardFp4QuantParams { - const void* __restrict__ q_input; - void* __restrict__ q_fp4; - int32_t* __restrict__ q_sf; - const void* __restrict__ weight; - float* __restrict__ weights_out; - float weight_scale; - const float* __restrict__ freqs_cis; - const void* __restrict__ positions; - uint32_t batch_size; - uint32_t num_heads; -}; - -template -Q_KERNEL void -fused_q_indexer_rope_hadamard_fp4_quant(const __grid_constant__ FusedQIndexerRopeHadamardFp4QuantParams params) { - using namespace device; - - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kVecSize = 4; - constexpr uint32_t kRopeSize = kRopeDim / kVecSize; - static_assert(kHeadDim == kWarpThreads * kVecSize); - static_assert(kRopeDim == kWarpThreads * 2); - static_assert(kRopeSize <= kWarpThreads); - - using Storage = AlignedVector; - using Float4 = AlignedVector; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto work_id = blockIdx.x * kFusedQNumWarps + warp_id; - const bool is_rope_lane = lane_id >= kWarpThreads - kRopeSize; - - const uint32_t total_works = params.batch_size * params.num_heads; - if (work_id >= total_works) return; - - const uint32_t batch_id = work_id / params.num_heads; - const auto input_ptr = static_cast(params.q_input) + work_id * kHeadDim; - const auto position = static_cast(static_cast(params.positions)[batch_id]); - const auto freqs_cis = params.freqs_cis + position * kRopeDim; - - PDLWaitPrimary(); - Float4 data, freq; - const auto weight_val = cast(static_cast(params.weight)[work_id]); - - { - Storage input_vec; - input_vec.load(input_ptr, lane_id); - if (is_rope_lane) freq.load(freqs_cis, lane_id - (kWarpThreads - kRopeSize)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - data[i] = cast(input_vec[i]); - } - } - - if (is_rope_lane) { - const auto x_real = data[0]; - const auto x_imag = data[1]; - const auto y_real = data[2]; - const auto y_imag = data[3]; - const auto fxr = freq[0]; - const auto fxi = freq[1]; - const auto fyr = freq[2]; - const auto fyi = freq[3]; - data[0] = x_real * fxr - x_imag * fxi; - data[1] = x_real * fxi + x_imag * fxr; - data[2] = y_real * fyr - y_imag * fyi; - data[3] = y_real * fyi + y_imag * fyr; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] = cast(cast(data[i])); - } - - PDLTriggerSecondary(); - - { - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a1; - data[1] = a0 - a1; - data[2] = a2 + a3; - data[3] = a2 - a3; - } - { - const float a0 = data[0], a1 = data[1], a2 = data[2], a3 = data[3]; - data[0] = a0 + a2; - data[1] = a1 + a3; - data[2] = a0 - a2; - data[3] = a1 - a3; - } -#pragma unroll - for (uint32_t mask = 1; mask < kWarpThreads; mask <<= 1) { -#pragma unroll - for (int i = 0; i < kVecSize; ++i) { - const float other = __shfl_xor_sync(0xFFFFFFFFu, data[i], mask, kWarpThreads); - data[i] = (lane_id & mask) ? (other - data[i]) : (data[i] + other); - } - } - const float kHadamardScale = math::rsqrt(static_cast(kHeadDim)); -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] *= kHadamardScale; -#pragma unroll - for (int i = 0; i < kVecSize; ++i) - data[i] = cast(cast(data[i])); - } - - { - float local_max = math::abs(data[0]); -#pragma unroll - for (int i = 1; i < kVecSize; ++i) { - local_max = math::max(local_max, math::abs(data[i])); - } - local_max = warp::reduce_max<8>(local_max); - const auto scale_raw = fmaxf(1e-4f, local_max) / 6.0f; - const auto scale_ue8m0 = static_cast(cast_to_ue8m0(scale_raw)); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); - const uint8_t packed0 = quant_fp4_e2m1(data[0] * inv_scale) | (quant_fp4_e2m1(data[1] * inv_scale) << 4); - const uint8_t packed1 = quant_fp4_e2m1(data[2] * inv_scale) | (quant_fp4_e2m1(data[3] * inv_scale) << 4); - const uint16_t packed = static_cast(packed0) | (static_cast(packed1) << 8); - auto out_row = static_cast(params.q_fp4) + work_id * (kHeadDim / 2); - reinterpret_cast(out_row)[lane_id] = packed; - if ((lane_id & 7) == 0) { - reinterpret_cast(params.q_sf + work_id)[lane_id >> 3] = scale_ue8m0; - } - params.weights_out[work_id] = weight_val * params.weight_scale; - } -} - -template -struct FusedQIndexerRopeHadamardFp4QuantKernel { - template - static constexpr auto kernel = fused_q_indexer_rope_hadamard_fp4_quant; - - static void forward( - const tvm::ffi::TensorView q_input, - const tvm::ffi::TensorView q_fp4, - const tvm::ffi::TensorView q_sf, - const tvm::ffi::TensorView weight, - const tvm::ffi::TensorView weights_out, - double weight_scale, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView positions) { - using namespace host; - constexpr int64_t kHeadDim = 128; - constexpr int64_t kRopeDim = 64; - constexpr int64_t kFp4Dim = kHeadDim / 2; - - auto B = SymbolicSize{"batch_size"}; - auto H = SymbolicSize{"num_heads"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, H, kHeadDim}) - .with_strides({-1, kHeadDim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_input); - TensorMatcher({B, H, kFp4Dim}) - .with_strides({-1, kFp4Dim, 1}) - .with_dtype() - .with_device(device_) - .verify(q_fp4); - TensorMatcher({B, H}).with_dtype().with_device(device_).verify(q_sf); - TensorMatcher({B, H}).with_dtype().with_device(device_).verify(weight); - TensorMatcher({B, H, 1}).with_dtype().with_device(device_).verify(weights_out); - TensorMatcher({-1, kRopeDim}).with_dtype().with_device(device_).verify(freqs_cis); - auto pos_dtype = SymbolicDType{}; - TensorMatcher({B}).with_dtype(pos_dtype).with_device(device_).verify(positions); - - const auto batch_size = static_cast(B.unwrap()); - const auto num_heads = static_cast(H.unwrap()); - if (batch_size == 0) return; - - const int64_t expected_q_stride = static_cast(num_heads) * kHeadDim; - const int64_t expected_fp4_stride = static_cast(num_heads) * kFp4Dim; - RuntimeCheck(q_input.stride(0) == expected_q_stride, "q_input must be contiguous"); - RuntimeCheck(q_fp4.stride(0) == expected_fp4_stride, "q_fp4 must be contiguous"); - RuntimeCheck(q_sf.stride(0) == static_cast(num_heads) && q_sf.stride(1) == 1, "q_sf must be contiguous"); - - const auto params = FusedQIndexerRopeHadamardFp4QuantParams{ - .q_input = q_input.data_ptr(), - .q_fp4 = q_fp4.data_ptr(), - .q_sf = static_cast(q_sf.data_ptr()), - .weight = weight.data_ptr(), - .weights_out = static_cast(weights_out.data_ptr()), - .weight_scale = static_cast(weight_scale), - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .positions = positions.data_ptr(), - .batch_size = batch_size, - .num_heads = num_heads, - }; - const auto total_works = batch_size * num_heads; - const auto num_blocks = div_ceil(total_works, kFusedQNumWarps); - const auto k_int32 = kernel; - const auto k_int64 = kernel; - const auto k = pos_dtype.is_type() ? k_int32 : k_int64; - LaunchKernel(num_blocks, kFusedQBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(k, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh deleted file mode 100644 index 7d5f97824b..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/mega_moe_pre_dispatch.cuh +++ /dev/null @@ -1,219 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include - -#include -#include - -namespace { - -using deepseek_v4::fp8::cast_to_ue8m0; -using deepseek_v4::fp8::pack_fp8; - -struct MegaMoEPreDispatchParams { - const bf16_t* __restrict__ x; // [num_tokens, hidden] - const int32_t* __restrict__ topk_idx; // [num_tokens, top_k] - const float* __restrict__ topk_weights; // [num_tokens, top_k] - - fp8_e4m3_t* __restrict__ buf_x; // [padded_max, hidden] - int32_t* __restrict__ buf_x_sf; // contiguous int32 [P, G/4]; see layout comment - int64_t* __restrict__ buf_topk_idx; // [padded_max, top_k] - float* __restrict__ buf_topk_weights; // [padded_max, top_k] - - uint32_t num_tokens; - uint32_t padded_max; - uint32_t hidden; - uint32_t num_groups; // hidden / group_size - uint32_t top_k; -}; - -// kGroupSize must match sglang_per_token_group_quant_fp8_ue8m0(group_size=). -template -__global__ __launch_bounds__(1024, 2) void // - mega_moe_pre_dispatch_kernel(const MegaMoEPreDispatchParams __grid_constant__ params) { - using namespace device; - - constexpr uint32_t kVecElems = 8; // 8 bf16 = 16B load per thread - static_assert(kGroupSize % kVecElems == 0, "group_size must be a multiple of 8"); - constexpr uint32_t kThreadsPerGroup = kGroupSize / kVecElems; - using InputVec = AlignedVector; - using OutputVec = AlignedVector; - - const uint32_t bid = blockIdx.x; - const uint32_t tid = threadIdx.x; - - PDLWaitPrimary(); - if (bid < params.num_tokens) { - // ---- Quantize path: one CTA per valid token ---- - - const uint32_t token_id = bid; - const auto token_in = params.x + static_cast(token_id) * params.hidden; - const auto token_out = params.buf_x + static_cast(token_id) * params.hidden; - - InputVec in_vec; - in_vec.load(token_in, tid); - - float local_max = 0.0f; - float vals[kVecElems]; -#pragma unroll - for (uint32_t i = 0; i < kVecElems / 2; ++i) { - const auto [v0, v1] = cast(in_vec[i]); - vals[2 * i + 0] = v0; - vals[2 * i + 1] = v1; - local_max = fmaxf(local_max, fmaxf(fabsf(v0), fabsf(v1))); - } - - // Absmax across the kThreadsPerGroup threads that cover one group. - local_max = warp::reduce_max(local_max); - - const float absmax = fmaxf(local_max, 1e-10f); - const float raw_scale = absmax / math::FP8_E4M3_MAX; - const uint32_t ue8m0_exp = cast_to_ue8m0(raw_scale); - // 2^-ue8m0_exp as fp32 (equivalent to 1 / __uint_as_float(ue8m0 << 23)). - const float inv_scale = __uint_as_float((127u + 127u - ue8m0_exp) << 23); - - OutputVec out_vec; -#pragma unroll - for (uint32_t i = 0; i < kVecElems / 2; ++i) { - out_vec[i] = pack_fp8(vals[2 * i + 0] * inv_scale, vals[2 * i + 1] * inv_scale); - } - out_vec.store(token_out, tid); - - // One thread per group writes its UE8M0 byte into the contiguous - // row-major int32-packed layout: byte address = t*num_groups + g - // (see layout comment at the top of the file). - const uint32_t group_id = tid / kThreadsPerGroup; - const uint32_t within_group_id = tid % kThreadsPerGroup; - if (within_group_id == 0 && group_id < params.num_groups) { - const uint32_t byte_off = token_id * params.num_groups + group_id; - reinterpret_cast(params.buf_x_sf)[byte_off] = static_cast(ue8m0_exp); - } - - // Copy this token's topk row (no alignment assumptions; top_k is small). - if (tid < params.top_k) { - const uint32_t off = token_id * params.top_k + tid; - params.buf_topk_idx[off] = params.topk_idx[off]; - params.buf_topk_weights[off] = params.topk_weights[off]; - } - } else { - // ---- Pad path: trailing blocks fill [num_tokens, padded_max) with (-1, 0) ---- - const uint32_t copy_bid = bid - params.num_tokens; - const uint32_t pad_base = params.num_tokens * params.top_k; - const uint32_t slot = pad_base + copy_bid * blockDim.x + tid; - const uint32_t total_slots = params.padded_max * params.top_k; - - if (slot < total_slots) { - params.buf_topk_idx[slot] = -1; - params.buf_topk_weights[slot] = 0.0f; - } - } - PDLTriggerSecondary(); -} - -// ---- Host wrapper -// ------------------------------------------------------------------------------------------------------------------------ - -template -struct MegaMoEPreDispatchKernel { - static_assert(kGroupSize == 32 || kGroupSize == 64 || kGroupSize == 128, "unsupported group_size"); - static constexpr auto kernel = mega_moe_pre_dispatch_kernel(kGroupSize), kUsePDL>; - - static void - run(const tvm::ffi::TensorView x, - const tvm::ffi::TensorView topk_idx, - const tvm::ffi::TensorView topk_weights, - const tvm::ffi::TensorView buf_x, - const tvm::ffi::TensorView buf_x_sf, - const tvm::ffi::TensorView buf_topk_idx, - const tvm::ffi::TensorView buf_topk_weights) { - using namespace host; - - auto device = SymbolicDevice{}; - auto M = SymbolicSize{"num_tokens"}; - auto P = SymbolicSize{"padded_max"}; - auto H = SymbolicSize{"hidden"}; - auto K = SymbolicSize{"top_k"}; - auto G4 = SymbolicSize{"num_groups_div_4"}; - device.set_options(); - - TensorMatcher({M, H}) // input x - .with_dtype() - .with_device(device) - .verify(x); - TensorMatcher({M, K}) // topk_idx - .with_dtype() - .with_device(device) - .verify(topk_idx); - TensorMatcher({M, K}) // topk_weights - .with_dtype() - .with_device(device) - .verify(topk_weights); - TensorMatcher({P, H}) // buf.x - .with_dtype() - .with_device(device) - .verify(buf_x); - // buf.x_sf is the contiguous row-major int32 view from DeepGEMM's mega - // symm buffer (DeepGEMM/csrc/apis/mega.hpp): shape (P, G/4), strides - // (G/4, 1). No explicit strides required -> TensorMatcher enforces - // is_contiguous(). - TensorMatcher({P, G4}) // buf_x_sf - .with_dtype() - .with_device(device) - .verify(buf_x_sf); - TensorMatcher({P, K}) // buf.topk_idx - .with_dtype() - .with_device(device) - .verify(buf_topk_idx); - TensorMatcher({P, K}) // buf.topk_weights - .with_dtype() - .with_device(device) - .verify(buf_topk_weights); - - const auto num_tokens = static_cast(M.unwrap()); - const auto padded_max = static_cast(P.unwrap()); - const auto hidden = static_cast(H.unwrap()); - const auto top_k = static_cast(K.unwrap()); - const auto num_groups_div_4 = static_cast(G4.unwrap()); - - RuntimeCheck(num_tokens <= padded_max, "num_tokens must not exceed padded_max"); - RuntimeCheck(hidden % kGroupSize == 0, "hidden must be a multiple of group_size"); - const auto num_groups = hidden / static_cast(kGroupSize); - RuntimeCheck(num_groups == num_groups_div_4 * 4u, "num_groups must be a multiple of 4"); - RuntimeCheck(hidden % 8u == 0, "hidden must be a multiple of 8 (16B bf16 loads)"); - const auto num_threads = hidden / 8u; - RuntimeCheck(num_threads <= 1024, "hidden too large for single-block-per-row quant"); - RuntimeCheck(num_threads >= top_k, "top_k must fit into one quant CTA"); - - const auto pad_slots = (padded_max - num_tokens) * top_k; - const uint32_t num_pad_blocks = pad_slots == 0 ? 0u : ((pad_slots + num_threads - 1u) / num_threads); - const auto num_total_blocks = num_tokens + num_pad_blocks; - - const auto params = MegaMoEPreDispatchParams{ - .x = static_cast(x.data_ptr()), - .topk_idx = static_cast(topk_idx.data_ptr()), - .topk_weights = static_cast(topk_weights.data_ptr()), - .buf_x = static_cast(buf_x.data_ptr()), - .buf_x_sf = static_cast(buf_x_sf.data_ptr()), - .buf_topk_idx = static_cast(buf_topk_idx.data_ptr()), - .buf_topk_weights = static_cast(buf_topk_weights.data_ptr()), - .num_tokens = num_tokens, - .padded_max = padded_max, - .hidden = hidden, - .num_groups = num_groups, - .top_k = top_k, - }; - - if (num_total_blocks == 0) return; - LaunchKernel(num_total_blocks, num_threads, device.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh deleted file mode 100644 index 38be975558..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/paged_mqa_metadata.cuh +++ /dev/null @@ -1,119 +0,0 @@ -#include -#include - -#include -#include - -#include -#include - -namespace { - -constexpr uint32_t kBlockSize = 1024; -constexpr uint32_t kSplitKV = 256; // const for both SM90 and SM100 - -struct MetadataParams { - /// NOTE: batch_size > 0 - uint32_t batch_size; - uint32_t num_sm; - const uint32_t* __restrict__ context_lens; - uint32_t* __restrict__ schedule_metadata; - bool use_smem = true; -}; - -__global__ __launch_bounds__(kBlockSize, 1) // - void smxx_paged_mqa_logits_metadata(const MetadataParams params) { - using namespace device; - extern __shared__ uint32_t s_length[]; - static constexpr auto kNumWarps = kBlockSize / kWarpThreads; - static_assert(kNumWarps == kWarpThreads); - - const auto tx = threadIdx.x; - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - __shared__ uint32_t s_warp_sum[kNumWarps]; - - uint32_t local_sum = 0; - for (uint32_t i = tx; i < params.batch_size; i += kBlockSize) { - const auto length = params.context_lens[i]; - local_sum += (length + kSplitKV - 1) / kSplitKV; - if (params.use_smem) s_length[i] = length; - } - - s_warp_sum[warp_id] = warp::reduce_sum(local_sum); - __syncthreads(); - - const auto global_sum = warp::reduce_sum(s_warp_sum[lane_id]); - if (lane_id != 0) return; - - const auto length_ptr = params.use_smem ? s_length : params.context_lens; - - const auto avg = global_sum / params.num_sm; - const auto ret = global_sum % params.num_sm; - uint32_t q = 0; - uint32_t num_work = (length_ptr[0] + kSplitKV - 1) / kSplitKV; - uint32_t sum_work = num_work; - for (auto i = warp_id; i <= params.num_sm; i += kNumWarps) { - const auto target = i * avg + min(i, ret); - while (sum_work <= target) { - if (++q >= params.batch_size) break; - num_work = (length_ptr[q] + kSplitKV - 1) / kSplitKV; - sum_work += num_work; - } - if (q >= params.batch_size) { - params.schedule_metadata[2 * i + 0] = params.batch_size; - params.schedule_metadata[2 * i + 1] = 0; - } else { - // sum > target && (sum - length) <= target - params.schedule_metadata[2 * i + 0] = q; - params.schedule_metadata[2 * i + 1] = target - (sum_work - num_work); - } - } -} - -template -void setup_kernel_smem_once(host::DebugInfo where = {}) { - [[maybe_unused]] - static const auto result = [] { - const auto fptr = std::bit_cast(f); - return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); - }(); - host::RuntimeDeviceCheck(result, where); -} - -struct IndexerMetadataKernel { - static constexpr auto kMaxBatchSizeInSmem = 16384 * 2; // 128 KB smeme - static void run(tvm::ffi::TensorView seq_lens, tvm::ffi::TensorView metadata) { - using namespace host; - auto N = SymbolicSize{"batch_size"}; - auto M = SymbolicSize{"num_sm"}; - auto device = SymbolicDevice{}; - device.set_options(); - TensorMatcher({N}) // - .with_dtype() - .with_device(device) - .verify(seq_lens); - TensorMatcher({M, 2}) // - .with_dtype() - .with_device(device) - .verify(metadata); - const auto batch_size = static_cast(N.unwrap()); - const auto num_sm = static_cast(M.unwrap()) - 1; - RuntimeCheck(num_sm <= 1024); - const auto use_smem = batch_size <= kMaxBatchSizeInSmem; - const auto params = MetadataParams{ - .batch_size = batch_size, - .num_sm = num_sm, - .context_lens = static_cast(seq_lens.data_ptr()), - .schedule_metadata = static_cast(metadata.data_ptr()), - .use_smem = use_smem, - }; - constexpr auto kernel = smxx_paged_mqa_logits_metadata; - setup_kernel_smem_once(); - const auto smem = use_smem ? (batch_size + 1) * sizeof(uint32_t) : 0; - LaunchKernel(1, kBlockSize, device.unwrap(), smem)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh deleted file mode 100644 index 2239d3972d..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/rope.cuh +++ /dev/null @@ -1,169 +0,0 @@ -#include -#include - -#include -#include -#include - -#include - -#include - -namespace { - -using DType = bf16_t; -constexpr int64_t kRopeDim = 64; -constexpr uint32_t kBlockSize = 128; -constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; - -struct FusedQKRopeParams { - void* __restrict__ q; - void* __restrict__ k; - const float* __restrict__ freqs_cis; - const void* __restrict__ positions; - int64_t q_stride_batch; - int64_t k_stride_batch; - int64_t q_stride_head; - int64_t k_stride_head; - uint32_t num_q_heads; - uint32_t num_k_heads; - uint32_t batch_size; -}; - -template -__global__ __launch_bounds__(kBlockSize, 16) // - void deepseek_rope_kernel(const __grid_constant__ FusedQKRopeParams param) { - using namespace device; - using DType2 = packed_t; - - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - const auto global_warp_id = blockIdx.x * kNumWarps + warp_id; - - const auto& [ - q, k, freqs_cis, positions, // - q_stride_batch, k_stride_batch, q_stride_head, k_stride_head, // - num_q_heads, num_k_heads, batch_size - ] = param; - - const auto num_total_heads = num_q_heads + num_k_heads; - const auto head_id = global_warp_id % num_total_heads; - const auto batch_id = global_warp_id / num_total_heads; - if (batch_id >= batch_size) return; - - const auto position = static_cast(positions)[batch_id]; - const auto is_q = head_id < num_q_heads; - const auto local_head = is_q ? head_id : (head_id - num_q_heads); - const auto stride_batch = is_q ? q_stride_batch : k_stride_batch; - const auto stride_head = is_q ? q_stride_head : k_stride_head; - const auto base_ptr = is_q ? q : k; - const auto input = static_cast(pointer::offset(base_ptr, batch_id * stride_batch, local_head * stride_head)); - - const auto freq_ptr = reinterpret_cast(freqs_cis + position * kRopeDim); - const auto [f_real, f_imag] = freq_ptr[lane_id]; - PDLWaitPrimary(); - - const auto data = input[lane_id]; - const auto [x_real, x_imag] = cast(data); - fp32x2_t output; - if constexpr (kInverse) { - // (a + bi) * (c - di) = (ac + bd) + (bc - ad)i - output = { - x_real * f_real + x_imag * f_imag, - x_imag * f_real - x_real * f_imag, - }; - } else { - // (a + bi) * (c + di) = (ac - bd) + (ad + bc)i - output = { - x_real * f_real - x_imag * f_imag, - x_real * f_imag + x_imag * f_real, - }; - } - input[lane_id] = cast(output); - - PDLTriggerSecondary(); -} - -template -struct FusedQKRopeKernel { - // 4 kernel variants: {forward, inverse} x {int32, int64} - static constexpr auto kernel_fwd_i32 = deepseek_rope_kernel; - static constexpr auto kernel_fwd_i64 = deepseek_rope_kernel; - static constexpr auto kernel_inv_i32 = deepseek_rope_kernel; - static constexpr auto kernel_inv_i64 = deepseek_rope_kernel; - - static void forward( - const tvm::ffi::TensorView q, - const tvm::ffi::Optional k, - const tvm::ffi::TensorView freqs_cis, - const tvm::ffi::TensorView positions, - bool inverse) { - using namespace host; - - auto B = SymbolicSize{"batch_size"}; - auto Q = SymbolicSize{"num_q_heads"}; - auto K = SymbolicSize{"num_k_heads"}; - constexpr auto D = kRopeDim; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, Q, D}) // - .with_strides({-1, -1, 1}) - .with_dtype() - .with_device(device_) - .verify(q); - if (k.has_value()) { - TensorMatcher({B, K, D}) // - .with_strides({-1, -1, 1}) - .with_dtype() - .with_device(device_) - .verify(k.value()); - } else { - K.set_value(0); - } - TensorMatcher({-1, D}) // - .with_dtype() - .with_device(device_) - .verify(freqs_cis); - - auto pos_dtype = SymbolicDType{}; - TensorMatcher({B}) // - .with_dtype(pos_dtype) - .with_device(device_) - .verify(positions); - const bool pos_i32 = pos_dtype.is_type(); - - const auto batch_size = static_cast(B.unwrap()); - if (batch_size == 0) return; - - const auto num_q_heads = static_cast(Q.unwrap()); - const auto num_k_heads = static_cast(K.unwrap()); - const auto num_total_heads = num_q_heads + num_k_heads; - const auto total_warps = batch_size * num_total_heads; - const auto num_blocks = div_ceil(total_warps, kNumWarps); - - const auto elem_size = static_cast(sizeof(DType)); - const auto params = FusedQKRopeParams{ - .q = q.data_ptr(), - .k = k ? k.value().data_ptr() : nullptr, - .freqs_cis = static_cast(freqs_cis.data_ptr()), - .positions = positions.data_ptr(), - .q_stride_batch = q.stride(0) * elem_size, - .k_stride_batch = k ? k.value().stride(0) * elem_size : 0, - .q_stride_head = q.stride(1) * elem_size, - .k_stride_head = k ? k.value().stride(1) * elem_size : 0, - .num_q_heads = num_q_heads, - .num_k_heads = num_k_heads, - .batch_size = batch_size, - }; - - // dispatch: {inverse} x {pos_i32} - using KernelType = decltype(kernel_fwd_i32); - const KernelType kernel = - inverse ? (pos_i32 ? kernel_inv_i32 : kernel_inv_i64) : (pos_i32 ? kernel_fwd_i32 : kernel_fwd_i64); - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh deleted file mode 100644 index be0e759445..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/silu_and_mul_masked_post_quant.cuh +++ /dev/null @@ -1,540 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include -#include - -#include - -#include -#include -#include - -namespace { - -using deepseek_v4::fp8::cast_to_ue8m0; -using deepseek_v4::fp8::pack_fp8; - -struct SiluMulQuantVarlenParams { - const bf16_t* __restrict__ input; - fp8_e4m3_t* __restrict__ output; - float* __restrict__ output_scale; - const int32_t* __restrict__ masked_m; - float swiglu_limit; // only read when kApplySwigluLimit=true - int64_t hidden_dim; - uint32_t num_tokens; - uint32_t num_experts; -}; - -constexpr uint32_t kMaxExperts = 256; - -struct alignas(16) CTAWork { - uint32_t expert_id; - uint32_t expert_token_id; - bool valid; -}; - -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(device::kWarpThreads == 32); -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { - uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); - if (lane_id >= offset) val += n; - } - return val; -} - -template -SGL_DEVICE fp32x2_t silu_and_mul(DType2 gate, DType2 up, float limit) { - using namespace device; - // refer to as implementation. TL;DR: must clamp in bf16 - // https://github.com/deepseek-ai/DeepGEMM/blob/7f2a703ed51ac1f7af07f5e1453b2d3267d37d50/deep_gemm/include/deep_gemm/impls/sm100_fp8_fp4_mega_moe.cuh#L984-L997 - if constexpr (kApplySwigluLimit) { - static_assert(std::is_same_v); - gate = __hmin2(gate, {limit, limit}); - up = __hmax2(up, {-limit, -limit}); - up = __hmin2(up, {limit, limit}); - } - const auto [g0, g1] = cast(gate); - const auto [u0, u1] = cast(up); - const auto silu0 = g0 / (1.0f + __expf(-g0)); - const auto silu1 = g1 / (1.0f + __expf(-g1)); - const float val0 = silu0 * u0; - const float val1 = silu1 * u1; - if constexpr (kPrecise) { // I don't know if we should enable this? - return {val0, val1}; - } else { - return cast(cast(fp32x2_t{val0, val1})); - } -} - -[[maybe_unused]] -SGL_DEVICE CTAWork get_work(const SiluMulQuantVarlenParams& params) { - // Preconditions: - // 1. blockDim.x >= params.num_experts - // 2. params.num_experts <= kMaxExperts - using namespace device; - static_assert(kWarpThreads == 32); - - static __shared__ uint32_t s_warp_sum[32]; - static __shared__ CTAWork result; - - result.valid = false; - - const uint32_t tx = threadIdx.x; - const uint32_t lane_id = tx % kWarpThreads; - const uint32_t warp_id = tx / kWarpThreads; - - const uint32_t val = tx < params.num_experts ? params.masked_m[tx] : 0u; - - // Per-warp inclusive scan of masked_m. - const uint32_t warp_inclusive = warp_inclusive_sum(lane_id, val); - const uint32_t warp_exclusive = warp_inclusive - val; - - // Write each warp total. - if (lane_id == kWarpThreads - 1) s_warp_sum[warp_id] = warp_inclusive; - __syncthreads(); - const auto tmp_val = lane_id < warp_id ? s_warp_sum[lane_id] : 0u; - const auto prefix_exclusive = warp::reduce_sum(tmp_val) + warp_exclusive; - const auto bx = blockIdx.x; - if (prefix_exclusive <= bx && bx < prefix_exclusive + val) { - result = {tx, bx - prefix_exclusive, true}; - } - __syncthreads(); - return result; -} - -template -__global__ __launch_bounds__(1024, 2) void // maximize occupancy - silu_mul_quant_varlen_kernel(const SiluMulQuantVarlenParams __grid_constant__ params) { - using namespace device; - - constexpr uint32_t kGroupSize = 128u; - constexpr uint32_t kWorkThreads = 16u; - // each thread will handle 8 elements - using InputVec = AlignedVector; - using OutputVec = AlignedVector; - static_assert(8 * kWorkThreads == 128, "Invalid tiling"); - static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0"); - - const auto [expert_id, token_id, valid] = get_work(params); - - if (!valid) return; - - const auto work_id = threadIdx.x / kWorkThreads; - - const auto offset = expert_id * params.num_tokens + token_id; - const auto input = params.input + offset * params.hidden_dim * 2; - const auto output = params.output + offset * params.hidden_dim; - [[maybe_unused]] - const auto output_scale = [&] { - const auto num_groups = params.hidden_dim / kGroupSize; - if constexpr (kTransposed) { - const auto base = reinterpret_cast(params.output_scale); - // Physical layout is [E, G//4, N] int32. Each int32 packs 4 consecutive - // group scales for the same token, so the byte address is: - // expert_offset + (group/4)*N*4 + token*4 + group%4 - return base + expert_id * num_groups * params.num_tokens + (work_id / 4u) * (params.num_tokens * 4u) + - token_id * 4u + (work_id % 4u); - } else { - return params.output_scale + offset * num_groups + work_id; - } - }(); - - PDLWaitPrimary(); - - InputVec gate_vec, up_vec; - if constexpr (kSwizzle) { - // gran=8 interleaved: every 16-element chunk on the N axis is - // [gate[0..7], up[0..7]]. Each thread handles 8 consecutive output - // elements, so its gate chunk lives at vec index 2*threadIdx.x and its - // up chunk at 2*threadIdx.x+1. - gate_vec.load(input, threadIdx.x * 2); - up_vec.load(input, threadIdx.x * 2 + 1); - } else { - gate_vec.load(input, threadIdx.x); - up_vec.load(input, threadIdx.x + blockDim.x); - } - - float local_max = 0.0f; - float results[8]; - -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const auto [x, y] = silu_and_mul(gate_vec[i], up_vec[i], params.swiglu_limit); - results[2 * i + 0] = x; - results[2 * i + 1] = y; - local_max = fmaxf(local_max, fmaxf(fabsf(x), fabsf(y))); - } - - local_max = warp::reduce_max(local_max); - - const float absmax = fmaxf(local_max, 1e-10f); - float scale; - uint32_t ue8m0_exp; - - if constexpr (kScaleUE8M0) { - const float raw_scale = absmax / math::FP8_E4M3_MAX; - ue8m0_exp = cast_to_ue8m0(raw_scale); - scale = __uint_as_float(ue8m0_exp << 23); - } else { - scale = absmax / math::FP8_E4M3_MAX; - } - const auto inv_scale = 1.0f / scale; - - OutputVec out_vec; -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const float scaled_val0 = results[2 * i + 0] * inv_scale; - const float scaled_val1 = results[2 * i + 1] * inv_scale; - out_vec[i] = pack_fp8(scaled_val0, scaled_val1); - } - - PDLTriggerSecondary(); - - out_vec.store(output, threadIdx.x); - if constexpr (kTransposed) { - *output_scale = ue8m0_exp; - } else { - *output_scale = scale; - } -} - -struct SiluAndMulClampParams { - const void* __restrict__ input; - void* __restrict__ output; - float swiglu_limit; -}; - -template -__global__ __launch_bounds__(1024, 2) void // maximize occupancy - silu_mul_clamp_kernel(const SiluAndMulClampParams __grid_constant__ params) { - using namespace device; - static_assert(sizeof(DType) == 2, "only fp16/bf16 supported"); - using DType2 = packed_t; - constexpr auto kVecSize = 16 / sizeof(DType); - static_assert(kVecSize % 2 == 0 && kVecSize > 0); - using Vec = AlignedVector; - const auto bid = blockIdx.x; - const auto tile = tile::Memory::cta(); - const float limit = params.swiglu_limit; - - PDLWaitPrimary(); - const auto gate = tile.load(params.input, bid * 2 + 0); - const auto up = tile.load(params.input, bid * 2 + 1); - Vec out; - -#pragma unroll - for (uint32_t i = 0; i < kVecSize / 2; ++i) { - out[i] = cast(silu_and_mul(cast(gate[i]), cast(up[i]), limit)); - } - - tile.store(params.output, out, bid); - PDLTriggerSecondary(); -} - -// ---- Host wrapper -// ------------------------------------------------------------------------------------------------------------------------ - -template -struct SiluAndMulMaskedPostQuantKernel { - static_assert(kGroupSize == 128); - static constexpr auto kernel_normal = - silu_mul_quant_varlen_kernel; - static constexpr auto kernel_transposed = - silu_mul_quant_varlen_kernel; - - static void - run(const tvm::ffi::TensorView input, - const tvm::ffi::TensorView output, - const tvm::ffi::TensorView output_scale, - const tvm::ffi::TensorView masked_m, - const uint32_t topk, - const bool transposed, - const double swiglu_limit) { - using namespace host; - - auto device = SymbolicDevice{}; - auto E = SymbolicSize{"num_experts"}; - auto T = SymbolicSize{"num_tokens_padded"}; - auto D = SymbolicSize{"hidden_dim x 2"}; - auto N = SymbolicSize{"hidden_dim"}; - auto G = SymbolicSize{"num_groups"}; - device.set_options(); - - TensorMatcher({E, T, D}) // input - .with_dtype() - .with_device(device) - .verify(input); - TensorMatcher({E, T, N}) // output - .with_dtype() - .with_device(device) - .verify(output); - if (!transposed) { - TensorMatcher({E, T, G}) // - .with_dtype() - .with_device(device) - .verify(output_scale); - } else { - RuntimeCheck(kScaleUE8M0, "transposed layout only supports scale_ue8m0=true"); - auto G_ = SymbolicSize{"G // 4"}; - TensorMatcher({E, G_, T}) // - .with_dtype() - .with_device(device) - .verify(output_scale); - G.set_value(G_.unwrap() * 4); - } - TensorMatcher({E}) // - .with_dtype() - .with_device(device) - .verify(masked_m); - - const auto num_experts = static_cast(E.unwrap()); - const auto num_tokens = static_cast(T.unwrap()); - const auto num_groups = static_cast(G.unwrap()); - const auto hidden_dim = N.unwrap(); - - RuntimeCheck(D.unwrap() == 2 * hidden_dim, "invalid dimension"); - RuntimeCheck(hidden_dim % kGroupSize == 0); - RuntimeCheck(num_experts <= kMaxExperts, "num_experts exceeds maximum (256)"); - RuntimeCheck(num_groups * kGroupSize == hidden_dim, "invalid num_groups"); - - const auto params = SiluMulQuantVarlenParams{ - .input = static_cast(input.data_ptr()), - .output = static_cast(output.data_ptr()), - .output_scale = static_cast(output_scale.data_ptr()), - .masked_m = static_cast(masked_m.data_ptr()), - .swiglu_limit = static_cast(swiglu_limit), - .hidden_dim = hidden_dim, - .num_tokens = num_tokens, - .num_experts = num_experts, - }; - - const auto num_threads = hidden_dim / 8; - RuntimeCheck(num_threads % device::kWarpThreads == 0); - RuntimeCheck(num_threads >= num_experts); - const auto kernel = transposed ? kernel_transposed : kernel_normal; - LaunchKernel(num_tokens * topk, num_threads, device.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -template -struct SiluAndMulClampKernel { - static constexpr auto kernel = silu_mul_clamp_kernel; - - static void run(const tvm::ffi::TensorView input, const tvm::ffi::TensorView output, const double swiglu_limit) { - using namespace host; - - auto device = SymbolicDevice{}; - auto M = SymbolicSize{"num_tokens"}; - auto D = SymbolicSize{"gate_up_dim"}; // 2 * out_dim - auto H = SymbolicSize{"out_dim"}; - device.set_options(); - - TensorMatcher({M, D}) // input (gate || up) - .with_dtype() - .with_device(device) - .verify(input); - TensorMatcher({M, H}) // output - .with_dtype() - .with_device(device) - .verify(output); - RuntimeCheck(D.unwrap() == 2 * H.unwrap(), "input last dim must be 2 * output last dim"); - - constexpr uint32_t kVecSize = 16 / sizeof(DType); - const auto out_dim = static_cast(H.unwrap()); - const auto num_tokens = static_cast(M.unwrap()); - RuntimeCheck(out_dim % kVecSize == 0, "out_dim must be divisible by vector size"); - const auto num_threads = out_dim / kVecSize; - RuntimeCheck(num_threads <= 1024, "out_dim too large for single-block-per-row launch"); - - const auto params = SiluAndMulClampParams{ - .input = input.data_ptr(), - .output = output.data_ptr(), - .swiglu_limit = static_cast(swiglu_limit), - }; - LaunchKernel(num_tokens, num_threads, device.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -struct SiluMulQuantContigParams { - const bf16_t* __restrict__ input; - fp8_e4m3_t* __restrict__ output; - float* __restrict__ output_scale; - float swiglu_limit; // only read when kApplySwigluLimit=true - int64_t hidden_dim; - uint32_t num_tokens; - uint32_t scale_row_stride_int32; // only used when kTransposed=true -}; - -template -__global__ __launch_bounds__(1024, 2) void // maximize occupancy - silu_mul_quant_contig_kernel(const SiluMulQuantContigParams __grid_constant__ params) { - using namespace device; - - constexpr uint32_t kGroupSize = 128u; - constexpr uint32_t kWorkThreads = 16u; - using InputVec = AlignedVector; - using OutputVec = AlignedVector; - static_assert(8 * kWorkThreads == 128, "Invalid tiling"); - static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0"); - - const auto token_id = blockIdx.x; - const auto work_id = threadIdx.x / kWorkThreads; - - const auto input = params.input + token_id * params.hidden_dim * 2; - const auto output = params.output + token_id * params.hidden_dim; - [[maybe_unused]] - const auto output_scale = [&] { - const auto num_groups = params.hidden_dim / kGroupSize; - if constexpr (kTransposed) { - // Physical layout is (G//4_pad, M_pad) int32; each int32 packs 4 - // consecutive UE8M0 exponents for the same token. Byte address: - // (work_id / 4) * M_pad * 4 + token * 4 + (work_id % 4). - const auto base = reinterpret_cast(params.output_scale); - return base + (work_id / 4u) * (params.scale_row_stride_int32 * 4u) + token_id * 4u + (work_id % 4u); - } else { - return params.output_scale + token_id * num_groups + work_id; - } - }(); - - PDLWaitPrimary(); - - InputVec gate_vec, up_vec; - if constexpr (kSwizzle) { - gate_vec.load(input, threadIdx.x * 2); - up_vec.load(input, threadIdx.x * 2 + 1); - } else { - gate_vec.load(input, threadIdx.x); - up_vec.load(input, threadIdx.x + blockDim.x); - } - - float local_max = 0.0f; - float results[8]; - -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const auto [x, y] = silu_and_mul(gate_vec[i], up_vec[i], params.swiglu_limit); - results[2 * i + 0] = x; - results[2 * i + 1] = y; - local_max = fmaxf(local_max, fmaxf(fabsf(x), fabsf(y))); - } - - local_max = warp::reduce_max(local_max); - - const float absmax = fmaxf(local_max, 1e-10f); - float scale; - uint32_t ue8m0_exp; - - if constexpr (kScaleUE8M0) { - const float raw_scale = absmax / math::FP8_E4M3_MAX; - ue8m0_exp = cast_to_ue8m0(raw_scale); - scale = __uint_as_float(ue8m0_exp << 23); - } else { - scale = absmax / math::FP8_E4M3_MAX; - } - const auto inv_scale = 1.0f / scale; - - OutputVec out_vec; -#pragma unroll - for (uint32_t i = 0; i < 4; ++i) { - const float scaled_val0 = results[2 * i + 0] * inv_scale; - const float scaled_val1 = results[2 * i + 1] * inv_scale; - out_vec[i] = pack_fp8(scaled_val0, scaled_val1); - } - - PDLTriggerSecondary(); - - out_vec.store(output, threadIdx.x); - if constexpr (kTransposed) { - *output_scale = ue8m0_exp; - } else { - *output_scale = scale; - } -} - -template -struct SiluAndMulContigPostQuantKernel { - static_assert(kGroupSize == 128); - static constexpr auto kernel_normal = - silu_mul_quant_contig_kernel; - static constexpr auto kernel_transposed = - silu_mul_quant_contig_kernel; - - static void - run(const tvm::ffi::TensorView input, - const tvm::ffi::TensorView output, - const tvm::ffi::TensorView output_scale, - const bool transposed, - const double swiglu_limit) { - using namespace host; - - auto device = SymbolicDevice{}; - auto M = SymbolicSize{"num_tokens"}; - auto D = SymbolicSize{"hidden_dim x 2"}; - auto N = SymbolicSize{"hidden_dim"}; - auto G = SymbolicSize{"num_groups"}; - device.set_options(); - - TensorMatcher({M, D}) // input (gate/up, natural or gran=8 interleaved on last dim) - .with_dtype() - .with_device(device) - .verify(input); - TensorMatcher({M, N}) // fp8 output - .with_dtype() - .with_device(device) - .verify(output); - - const auto hidden_dim = N.unwrap(); - RuntimeCheck(D.unwrap() == 2 * hidden_dim, "invalid dimension"); - RuntimeCheck(hidden_dim % kGroupSize == 0); - const auto num_groups = static_cast(hidden_dim / kGroupSize); - - uint32_t scale_row_stride_int32 = 0; - if (!transposed) { - G.set_value(num_groups); - TensorMatcher({M, G}) // (M, G) fp32 natural row-major - .with_dtype() - .with_device(device) - .verify(output_scale); - } else { - RuntimeCheck(kScaleUE8M0, "transposed layout only supports scale_ue8m0=true"); - RuntimeCheck(num_groups % 4 == 0, "transposed layout requires num_groups % 4 == 0"); - auto G_ = SymbolicSize{"G // 4"}; - G_.set_value(num_groups / 4); - auto M_pad = SymbolicSize{"M padded"}; - TensorMatcher({M, G_}) // `.transpose(-1,-2)[:M,:]` view of (G//4_pad, M_pad) int32 - .with_strides({int64_t{1}, M_pad}) // col-major transposed - .with_dtype() - .with_device(device) - .verify(output_scale); - scale_row_stride_int32 = static_cast(M_pad.unwrap()); - } - - const auto num_tokens = static_cast(M.unwrap()); - - const auto params = SiluMulQuantContigParams{ - .input = static_cast(input.data_ptr()), - .output = static_cast(output.data_ptr()), - .output_scale = static_cast(output_scale.data_ptr()), - .swiglu_limit = static_cast(swiglu_limit), - .hidden_dim = hidden_dim, - .num_tokens = num_tokens, - .scale_row_stride_int32 = scale_row_stride_int32, - }; - - const auto num_threads = hidden_dim / 8; - RuntimeCheck(num_threads % device::kWarpThreads == 0); - const auto kernel = transposed ? kernel_transposed : kernel_normal; - LaunchKernel(num_tokens, num_threads, device.unwrap()) // - .enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh deleted file mode 100644 index 49f6f55963..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/store.cuh +++ /dev/null @@ -1,205 +0,0 @@ -#include -#include - -#include -#include -#include -#include -#include - -#include - -#include -#include - -#include -#include -#include - -namespace { - -using deepseek_v4::fp8::cast_to_ue8m0; -using deepseek_v4::fp8::inv_scale_ue8m0; -using deepseek_v4::fp8::pack_fp8; - -struct FusedStoreCacheParam { - const void* __restrict__ input; - void* __restrict__ cache; - const void* __restrict__ indices; - uint32_t num_tokens; -}; - -template -__global__ void fused_store_flashmla_cache(const __grid_constant__ FusedStoreCacheParam param) { - using namespace device; - - /// NOTE: 584 = 576 + 8 - constexpr int64_t kPageBytes = host::div_ceil(584 << kPageBits, 576) * 576; - - // each warp handles 64 elements, 8 warps, each block handles 1 row - const auto& [input, cache, indices, num_tokens] = param; - const uint32_t bid = blockIdx.x; - const uint32_t tid = threadIdx.x; - const uint32_t wid = tid / 32; - - PDLWaitPrimary(); - - // prefetch the index - const auto index = static_cast(indices)[bid]; - // always load the value from input (don't store if invalid) - using Float2 = packed_t; - const auto elems = static_cast(input)[tid + bid * 256]; - if (wid != 7) { - const auto [x, y] = cast(elems); - const auto abs_max = warp::reduce_max(fmaxf(fabs(x), fabs(y))); - const auto scale_raw = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto scale_ue8m0 = cast_to_ue8m0(scale_raw); - const auto inv_scale = inv_scale_ue8m0(scale_ue8m0); - const auto result = pack_fp8(x * inv_scale, y * inv_scale); - const int32_t page = index >> kPageBits; - const int32_t offset = index & ((1 << kPageBits) - 1); - const auto page_ptr = pointer::offset(cache, page * kPageBytes); - const auto value_ptr = pointer::offset(page_ptr, offset * 576); - const auto scale_ptr = pointer::offset(page_ptr, 576 << kPageBits, offset * 8); - static_cast(value_ptr)[tid] = result; - static_cast(scale_ptr)[wid] = scale_ue8m0; - } else { - const auto result = cast(elems); - const int32_t page = index >> kPageBits; - const int32_t offset = index & ((1 << kPageBits) - 1); - const auto page_ptr = pointer::offset(cache, page * kPageBytes); - const auto value_ptr = pointer::offset(page_ptr, offset * 576, 448); - static_cast(value_ptr)[tid - 7 * 32] = result; - } - - PDLTriggerSecondary(); -} - -template -__global__ void fused_store_indexer_cache(const __grid_constant__ FusedStoreCacheParam param) { - using namespace device; - - /// NOTE: 132 = 128 + 4 - constexpr int64_t kPageBytes = 132 << kPageBits; - - // each warp handles 128 elements, 1 warp, each block handles multiple rows - const auto& [input, cache, indices, num_tokens] = param; - const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; - const auto global_wid = global_tid / 32; - const auto lane_id = threadIdx.x % 32; - - if (global_wid >= num_tokens) return; - - PDLWaitPrimary(); - - // prefetch the index - const auto index = static_cast(indices)[global_wid]; - // always load the value from input (don't store if invalid) - using Float2 = packed_t; - using InStorage = AlignedVector; - using OutStorage = AlignedVector; - const auto elems = static_cast(input)[global_tid]; - const auto [x0, x1] = cast(elems[0]); - const auto [y0, y1] = cast(elems[1]); - const auto local_max = fmaxf(fmaxf(fabs(x0), fabs(x1)), fmaxf(fabs(y0), fabs(y1))); - const auto abs_max = warp::reduce_max(local_max); - // use normal fp32 scale - const auto scale = fmaxf(1e-4f, abs_max) / math::FP8_E4M3_MAX; - const auto inv_scale = 1.0f / scale; - const int32_t page = index >> kPageBits; - const int32_t offset = index & ((1 << kPageBits) - 1); - const auto page_ptr = pointer::offset(cache, page * kPageBytes); - const auto value_ptr = pointer::offset(page_ptr, offset * 128); - const auto scale_ptr = pointer::offset(page_ptr, 128 << kPageBits, offset * 4); - OutStorage result; - result[0] = pack_fp8(x0 * inv_scale, x1 * inv_scale); - result[1] = pack_fp8(y0 * inv_scale, y1 * inv_scale); - static_cast(value_ptr)[lane_id] = result; - static_cast(scale_ptr)[0] = scale; - - PDLTriggerSecondary(); -} - -template -struct FusedStoreCacheFlashMLAKernel { - static constexpr int32_t kLogSize = std::countr_zero(kPageSize); - static constexpr int64_t kPageBytes = host::div_ceil(584 * kPageSize, 576) * 576; - static constexpr auto kernel = fused_store_flashmla_cache; - - static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); - static_assert(1 << kLogSize == kPageSize); - - static void run(tvm::ffi::TensorView input, tvm::ffi::TensorView cache, tvm::ffi::TensorView indices) { - using namespace host; - - auto N = SymbolicSize{"num_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - TensorMatcher({N, 512}) // input - .with_dtype() - .with_device(device_) - .verify(input); - TensorMatcher({-1, -1}) // cache - .with_strides({kPageBytes, 1}) - .with_dtype() - .with_device(device_) - .verify(cache); - TensorMatcher({N}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - const auto num_tokens = static_cast(N.unwrap()); - const auto params = FusedStoreCacheParam{ - .input = input.data_ptr(), - .cache = cache.data_ptr(), - .indices = indices.data_ptr(), - .num_tokens = num_tokens, - }; - const auto kBlockSize = 256; - const auto num_blocks = num_tokens; - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); - } -}; - -template -struct FusedStoreCacheIndexerKernel { - static constexpr int32_t kLogSize = std::countr_zero(kPageSize); - static constexpr int64_t kPageBytes = 132 * kPageSize; - static constexpr auto kernel = fused_store_indexer_cache; - - static_assert(std::has_single_bit(kPageSize), "kPageSize must be a power of 2"); - static_assert(1 << kLogSize == kPageSize); - - static void run(tvm::ffi::TensorView input, tvm::ffi::TensorView cache, tvm::ffi::TensorView indices) { - using namespace host; - - auto N = SymbolicSize{"num_tokens"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - TensorMatcher({N, 128}) // input - .with_dtype() - .with_device(device_) - .verify(input); - TensorMatcher({-1, -1}) // cache - .with_strides({kPageBytes, 1}) - .with_dtype() - .with_device(device_) - .verify(cache); - TensorMatcher({N}) // indices - .with_dtype() - .with_device(device_) - .verify(indices); - const auto num_tokens = static_cast(N.unwrap()); - const auto params = FusedStoreCacheParam{ - .input = input.data_ptr(), - .cache = cache.data_ptr(), - .indices = indices.data_ptr(), - .num_tokens = num_tokens, - }; - const auto kBlockSize = 128; - const auto num_blocks = div_ceil(num_tokens * 32, kBlockSize); - LaunchKernel(num_blocks, kBlockSize, device_.unwrap()).enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh deleted file mode 100644 index b1ccd24b20..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v1.cuh +++ /dev/null @@ -1,340 +0,0 @@ -#include -#include - -#include - -#include -#include - -#include -#include - -namespace { - -#ifndef SGL_TOPK -#define SGL_TOPK 512 -#endif - -constexpr uint32_t kTopK = SGL_TOPK; -constexpr uint32_t kTopKBlockSize = SGL_TOPK; -constexpr uint32_t kSMEM = 16 * 1024 * sizeof(uint32_t); // 64KB (bytes) - -struct TopKParams { - const float* __restrict__ scores; - const int32_t* __restrict__ seq_lens; - const int32_t* __restrict__ page_table; - int32_t* __restrict__ page_indices; - int32_t* __restrict__ raw_indices; // optional: output raw abs position indices before page transform - const int64_t score_stride; - const int64_t page_table_stride; - uint32_t page_bits; -}; - -SGL_DEVICE uint8_t convert_to_uint8(float x) { - __half h = __float2half_rn(x); - uint16_t bits = __half_as_ushort(h); - uint16_t key = (bits & 0x8000) ? static_cast(~bits) : static_cast(bits | 0x8000); - return static_cast(key >> 8); -} - -SGL_DEVICE uint32_t convert_to_uint32(float x) { - uint32_t bits = __float_as_uint(x); - return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); -} - -SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) { - const uint32_t mask = (1u << page_bits) - 1u; - return (page_table[i >> page_bits] << page_bits) | (i & mask); -} - -[[maybe_unused]] -SGL_DEVICE void naive_transform( - const float* __restrict__, // unused - const int32_t* __restrict__ page_table, - int32_t* __restrict__ indices, - int32_t* __restrict__ raw_indices, // optional: output raw abs position indices - const uint32_t length, - const uint32_t page_bits) { - static_assert(kTopK <= kTopKBlockSize); - if (const auto tx = threadIdx.x; tx < length) { - indices[tx] = page_to_indices(page_table, tx, page_bits); - if (raw_indices != nullptr) { - raw_indices[tx] = tx; - } - } else if (kTopK == kTopKBlockSize || tx < kTopK) { - indices[tx] = -1; // fill invalid indices to -1 - if (raw_indices != nullptr) { - raw_indices[tx] = -1; - } - } -} - -[[maybe_unused]] -SGL_DEVICE void radix_topk(const float* __restrict__ input, int32_t* __restrict__ output, const uint32_t length) { - constexpr uint32_t RADIX = 256; - constexpr uint32_t BLOCK_SIZE = kTopKBlockSize; - constexpr uint32_t SMEM_INPUT_SIZE = kSMEM / (2 * sizeof(int32_t)); - - alignas(128) __shared__ uint32_t _s_histogram_buf[2][RADIX + 32]; - alignas(128) __shared__ uint32_t s_counter; - alignas(128) __shared__ uint32_t s_threshold_bin_id; - alignas(128) __shared__ uint32_t s_num_input[2]; - alignas(128) __shared__ int32_t s_last_remain; - - extern __shared__ uint32_t s_input_idx[][kSMEM / (2 * sizeof(int32_t))]; - - const uint32_t tx = threadIdx.x; - uint32_t remain_topk = kTopK; - auto& s_histogram = _s_histogram_buf[0]; - - const auto run_cumsum = [&] { -#pragma unroll 8 - for (int32_t i = 0; i < 8; ++i) { - static_assert(1 << 8 == RADIX); - if (tx < RADIX) { - const auto j = 1 << i; - const auto k = i & 1; - auto value = _s_histogram_buf[k][tx]; - if (tx + j < RADIX) { - value += _s_histogram_buf[k][tx + j]; - } - _s_histogram_buf[k ^ 1][tx] = value; - } - __syncthreads(); - } - }; - - // stage 1: 8bit coarse histogram - if (tx < RADIX + 1) s_histogram[tx] = 0; - __syncthreads(); - for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { - const auto bin = convert_to_uint8(input[idx]); - ::atomicAdd(&s_histogram[bin], 1); - } - __syncthreads(); - run_cumsum(); - if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { - s_threshold_bin_id = tx; - s_num_input[0] = 0; - s_counter = 0; - } - __syncthreads(); - - const auto threshold_bin = s_threshold_bin_id; - remain_topk -= s_histogram[threshold_bin + 1]; - if (remain_topk == 0) { - for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { - const uint32_t bin = convert_to_uint8(input[idx]); - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - output[pos] = idx; - } - } - __syncthreads(); - return; - } else { - __syncthreads(); - if (tx < RADIX + 1) { - s_histogram[tx] = 0; - } - __syncthreads(); - - for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) { - const float raw_input = input[idx]; - const uint32_t bin = convert_to_uint8(raw_input); - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - output[pos] = idx; - } else if (bin == threshold_bin) { - const auto pos = ::atomicAdd(&s_num_input[0], 1); - if (pos < SMEM_INPUT_SIZE) { - [[likely]] s_input_idx[0][pos] = idx; - const auto bin = convert_to_uint32(raw_input); - const auto sub_bin = (bin >> 24) & 0xFF; - ::atomicAdd(&s_histogram[sub_bin], 1); - } - } - } - __syncthreads(); - } - - // stage 2: refine with 8bit radix passes -#pragma unroll 4 - for (int round = 0; round < 4; ++round) { - const auto r_idx = round % 2; - - // clip here to prevent overflow - const auto raw_num_input = s_num_input[r_idx]; - const auto num_input = raw_num_input < SMEM_INPUT_SIZE ? raw_num_input : SMEM_INPUT_SIZE; - - run_cumsum(); - if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) { - s_threshold_bin_id = tx; - s_num_input[r_idx ^ 1] = 0; - s_last_remain = remain_topk - s_histogram[tx + 1]; - } - __syncthreads(); - - const auto threshold_bin = s_threshold_bin_id; - remain_topk -= s_histogram[threshold_bin + 1]; - - if (remain_topk == 0) { - for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { - const auto idx = s_input_idx[r_idx][i]; - const auto offset = 24 - round * 8; - const auto bin = (convert_to_uint32(input[idx]) >> offset) & 0xFF; - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - output[pos] = idx; - } - } - __syncthreads(); - break; - } else { - __syncthreads(); - if (tx < RADIX + 1) { - s_histogram[tx] = 0; - } - __syncthreads(); - for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) { - const auto idx = s_input_idx[r_idx][i]; - const auto raw_input = input[idx]; - const auto offset = 24 - round * 8; - const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF; - if (bin > threshold_bin) { - const auto pos = ::atomicAdd(&s_counter, 1); - output[pos] = idx; - } else if (bin == threshold_bin) { - if (round == 3) { - const auto pos = ::atomicAdd(&s_last_remain, -1); - if (pos > 0) { - output[kTopK - pos] = idx; - } - } else { - const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1); - if (pos < SMEM_INPUT_SIZE) { - /// NOTE: (dark) fuse the histogram computation here - [[likely]] s_input_idx[r_idx ^ 1][pos] = idx; - const auto bin = convert_to_uint32(raw_input); - const auto sub_bin = (bin >> (offset - 8)) & 0xFF; - ::atomicAdd(&s_histogram[sub_bin], 1); - } - } - } - } - __syncthreads(); - } - } -} - -template -__global__ void topk_transform_kernel(const __grid_constant__ TopKParams params) { - const auto &[ - scores, seq_lens, page_table, page_indices, raw_indices, // pointers - score_stride, page_table_stride, page_bits // sizes - ] = params; - const uint32_t work_id = blockIdx.x; - - /// NOTE: dangerous prefetch seq_len before PDL wait - const uint32_t seq_len = seq_lens[work_id]; - const auto score_ptr = scores + work_id * score_stride; - const auto page_ptr = page_table + work_id * page_table_stride; - const auto indices_ptr = page_indices + work_id * kTopK; - const auto raw_indices_ptr = raw_indices != nullptr ? raw_indices + work_id * kTopK : nullptr; - - device::PDLWaitPrimary(); - - if (seq_len <= kTopK) { - naive_transform(score_ptr, page_ptr, indices_ptr, raw_indices_ptr, seq_len, page_bits); - } else { - __shared__ int32_t s_topk_indices[kTopK]; - radix_topk(score_ptr, s_topk_indices, seq_len); - static_assert(kTopK <= kTopKBlockSize); - const auto tx = threadIdx.x; - if (kTopK == kTopKBlockSize || tx < kTopK) { - indices_ptr[tx] = page_to_indices(page_ptr, s_topk_indices[tx], page_bits); - if (raw_indices_ptr != nullptr) { - raw_indices_ptr[tx] = s_topk_indices[tx]; - } - } - } - - device::PDLTriggerSecondary(); -} - -template -void setup_kernel_smem_once(host::DebugInfo where = {}) { - [[maybe_unused]] - static const auto result = [] { - const auto fptr = std::bit_cast(f); - return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); - }(); - host::RuntimeDeviceCheck(result, where); -} - -template -struct TopKKernel { - static constexpr auto kernel = topk_transform_kernel; - - static void transform( - const tvm::ffi::TensorView scores, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView page_table, - const tvm::ffi::TensorView page_indices, - const uint32_t page_size, - const tvm::ffi::Optional raw_indices) { - using namespace host; - auto B = SymbolicSize{"batch_size"}; - auto S = SymbolicSize{"score_stride"}; - auto P = SymbolicSize{"page_table_stride"}; - auto device = SymbolicDevice{}; - device.set_options(); - - TensorMatcher({B, -1}) // strided scores - .with_strides({S, 1}) - .with_dtype() - .with_device(device) - .verify(scores); - TensorMatcher({B}) // seq_lens, must be contiguous - .with_dtype() - .with_device(device) - .verify(seq_lens); - TensorMatcher({B, -1}) // strided page table - .with_strides({P, 1}) - .with_dtype() - .with_device(device) - .verify(page_table); - TensorMatcher({B, kTopK}) // output, must be contiguous - .with_dtype() - .with_device(device) - .verify(page_indices); - - int32_t* raw_indices_ptr = nullptr; - if (raw_indices.has_value()) { - TensorMatcher({B, kTopK}) // optional raw indices output, must be contiguous - .with_dtype() - .with_device(device) - .verify(raw_indices.value()); - raw_indices_ptr = static_cast(raw_indices.value().data_ptr()); - } - - RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2"); - const auto page_bits = static_cast(std::countr_zero(page_size)); - const auto batch_size = static_cast(B.unwrap()); - const auto params = TopKParams{ - .scores = static_cast(scores.data_ptr()), - .seq_lens = static_cast(seq_lens.data_ptr()), - .page_table = static_cast(page_table.data_ptr()), - .page_indices = static_cast(page_indices.data_ptr()), - .raw_indices = raw_indices_ptr, - .score_stride = S.unwrap(), - .page_table_stride = P.unwrap(), - .page_bits = page_bits, - }; - constexpr auto kSMEM_ = kSMEM + sizeof(int32_t); // align up a little - setup_kernel_smem_once(); - LaunchKernel(batch_size, kTopKBlockSize, device.unwrap(), kSMEM_).enable_pdl(kUsePDL)(kernel, params); - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh b/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh deleted file mode 100644 index 8c4a526575..0000000000 --- a/lightllm/third_party/sglang_jit/csrc/deepseek_v4/topk_v2.cuh +++ /dev/null @@ -1,493 +0,0 @@ -#include -#include - -#include -#include -#include -#include - -#include -#include -#include - -#include -#include -#include - -#include -#include -#include - -namespace { - -#ifndef SGL_TOPK -#define SGL_TOPK 512 -#endif - -inline constexpr uint32_t K = SGL_TOPK; - -template -void setup_kernel_smem_once(host::DebugInfo where = {}) { - [[maybe_unused]] - static const auto result = [] { - const auto fptr = std::bit_cast(f); - return ::cudaFuncSetAttribute(fptr, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM); - }(); - host::RuntimeDeviceCheck(result, where); -} - -namespace impl = device::top512; -using Large = impl::ClusterTopK; -using Medium = impl::StreamingTopK; -using Small = impl::RegisterTopK; - -using Metadata = Large::Metadata; -constexpr uint32_t kBlockSize = impl::kBlockSize; -constexpr uint32_t kNumClusters = 15; // based on hardware limits -constexpr uint32_t kClusterSize = Large::kClusterSize; -constexpr uint32_t kMax2PassLength = Small::kMax2PassLength; -constexpr uint32_t kMaxSupportedLength = Large::kMaxLength; - -/// Common metadata lives at metadata[0] (first row of the [batch_size+1, 4] tensor). -/// Per-item metadata starts at metadata[1..batch_size]. The plan kernel writes both. -struct alignas(16) GlobalMetadata { - uint32_t cluster_threshold; // decided per-batch in plan kernel - uint32_t num_cluster_items; // N = number of items routed to the cluster path - uint32_t reserved[2]; -}; -static_assert(sizeof(GlobalMetadata) == sizeof(Metadata), "layout: row 0 must occupy one Metadata-sized slot"); - -// optimize occupancy for prefill -#define SMALL_TOPK_KERNEL __global__ __launch_bounds__(kBlockSize, 2) -// cluster at y dim -#define LARGE_CLUSTER __cluster_dims__(1, kClusterSize, 1) -// stage-1 is persistent cluster, and shared memory usage is huge (can not 2) -#define LARGE_TOPK_STAGE_1 __global__ __launch_bounds__(kBlockSize, 1) LARGE_CLUSTER -// stage-2 is non-persistent non-cluster, with less shared memory and higher occupancy -#define LARGE_TOPK_STAGE_2 __global__ __launch_bounds__(kBlockSize, 2) -// fused into 1 stage when batch-size <= kNumPersistentClusters -#define FUSED_COMBINE_KERNEL __global__ __launch_bounds__(kBlockSize, 1) LARGE_CLUSTER -// plan runs once as a single block before the combine kernels -#define PLAN_KERNEL __global__ __launch_bounds__(kBlockSize, 1) - -struct TopKParams { - const uint32_t* __restrict__ seq_lens; - const float* __restrict__ scores; - const int32_t* __restrict__ page_table; - int32_t* __restrict__ page_indices; - int64_t score_stride; - int64_t page_table_stride; - uint8_t* __restrict__ workspace; // [batch, kWorkspaceBytes] -- internally allocated - /// Pointer to the full metadata tensor: metadata[0] is GlobalMetadata, metadata[1..] - /// are per-item entries (at most kNumClusters * rounds of them). - const Metadata* __restrict__ metadata = nullptr; - int64_t workspace_stride; // bytes per batch - uint32_t batch_size; - uint32_t page_bits; - - SGL_DEVICE const float* get_scores(const uint32_t batch_id) const { - return scores + batch_id * score_stride; - } - SGL_DEVICE impl::TransformParams get_transform(const uint32_t batch_id, int32_t* indices) const { - return { - .page_table = page_table + batch_id * page_table_stride, - .indices_in = indices, - .indices_out = page_indices + batch_id * K, - .page_bits = page_bits, - }; - } - SGL_DEVICE const GlobalMetadata& get_global_metadata() const { - return *reinterpret_cast(metadata); - } - SGL_DEVICE const Metadata& get_item_metadata(uint32_t work_id) const { - return metadata[1 + work_id]; // +1 to skip the GlobalMetadata row - } -}; - -SGL_DEVICE uint2 partition_work(uint32_t length, uint32_t rank) { - constexpr uint32_t kTMAAlign = 4; - const auto total_units = (length + kTMAAlign - 1) / kTMAAlign; - const auto base = total_units / kClusterSize; - const auto extra = total_units % kClusterSize; - const auto local_units = base + (rank < extra ? 1u : 0u); - const auto offset_units = rank * base + min(rank, extra); - const auto offset = offset_units * kTMAAlign; - const auto finish = min(offset + local_units * kTMAAlign, length); - return {offset, finish - offset}; -} - -/// Persistent scheduler. A single block: -/// 1. Decides a cluster_threshold from the real seq_lens distribution (or -/// uses the caller-supplied `static_cluster_threshold` when non-zero). -/// 2. Writes that threshold + N into metadata[0] (the GlobalMetadata row). -/// 3. Compacts items with seq_len > threshold into metadata[1..N+1), laid out -/// to match the persistent consumer's round-robin stride (kNumClusters). -/// Entries for clusters that get no work are zero-filled. -PLAN_KERNEL void topk_plan( - const uint32_t* __restrict__ seq_lens, - Metadata* __restrict__ metadata, - const uint32_t batch_size, - const uint32_t static_cluster_threshold) { - // Candidate thresholds, strictly increasing. Picked to give the auto-heuristic - // reasonable granularity without needing a full sort. Must all be >= kMax2PassLength. - - struct Pair { - uint32_t threshold; - uint32_t max_batch_size; - }; - /// NOTE: only tuned on B200 - constexpr Pair kCandidates[] = { - {32768, 30}, - {40960, 45}, - {49152, 45}, - {65536, 60}, - {98304, 60}, - {131072, 75}, - {196608, 90}, - {262144, 105}, - }; - constexpr uint32_t kNumCandidates = std::size(kCandidates); - constexpr uint32_t kMinBatchSize = kCandidates[0].max_batch_size; - static_assert(kCandidates[0].threshold == kMax2PassLength); - static_assert(kCandidates[kNumCandidates - 1].threshold == kMaxSupportedLength); - - __shared__ uint32_t s_count; // final N after compaction - __shared__ uint32_t s_counts[kNumCandidates]; - __shared__ uint32_t s_threshold; - - const auto tx = threadIdx.x; - if (tx == 0) s_count = 0; - if (tx < kNumCandidates) s_counts[tx] = 0; - __syncthreads(); - - // --- Phase 1: decide threshold ------------------------------------------ - if (static_cluster_threshold > 0) { - if (tx == 0) s_threshold = static_cluster_threshold; - } else if (batch_size <= kMinBatchSize) { - if (tx == 0) s_threshold = kMax2PassLength; // always prefer cluster - } else { - // Count items above each candidate threshold. Monotonically non-increasing in T. - for (uint32_t i = tx; i < batch_size; i += kBlockSize) { - const uint32_t sl = seq_lens[i]; - assert(sl <= kMaxSupportedLength); - uint32_t count = 0; -#pragma unroll - for (uint32_t j = 0; j < kNumCandidates; ++j) { - count += (sl > kCandidates[j].threshold ? 1 : 0); - } - if (count > 0) { - atomicAdd(&s_counts[count - 1], 1); - } - } - __syncthreads(); - if (tx == 0) { - uint32_t accum = 0; - uint32_t chosen = kMaxSupportedLength; -#pragma unroll - for (uint32_t i = 0; i < kNumCandidates; ++i) { - const auto j = kNumCandidates - 1 - i; - accum += s_counts[j]; - /// NOTE: `accum` increasing, while `max_batch_size` decreasing - if (accum > kCandidates[j].max_batch_size) break; - chosen = kCandidates[j].threshold; - } - s_threshold = chosen; - } - } - __syncthreads(); - // sanity check: below 2 pass threshold, must fits in small path - const auto cluster_threshold = max(s_threshold, kMax2PassLength); - - // --- Phase 2: compact items with seq_len > threshold into metadata[1..] - - // Per-item rows live at metadata[1 + pos]; metadata[0] is the GlobalMetadata row. - for (uint32_t i = tx; i < batch_size; i += kBlockSize) { - const uint32_t sl = seq_lens[i]; - if (sl > cluster_threshold) { - const auto pos = atomicAdd(&s_count, 1); - metadata[1 + pos] = {i, sl, false}; - } - } - __syncthreads(); - const auto N = s_count; - - // --- Phase 3: has_next + sentinels + GlobalMetadata --------------------- - for (uint32_t i = tx; i < N; i += kBlockSize) { - if (i + kNumClusters < N) metadata[1 + i].has_next = true; - } - // Zero-fill the first kNumClusters sentinel slots that got no valid entry. - if (tx < kNumClusters && tx >= N) metadata[1 + tx] = {0, 0, false}; - // Write global metadata (row 0). - if (tx == 0) { - auto* g = reinterpret_cast(metadata); - *g = { - .cluster_threshold = cluster_threshold, - .num_cluster_items = N, - .reserved = {0, 0}, - }; - } -} - -SMALL_TOPK_KERNEL void // short context -topk_short_transform(const __grid_constant__ TopKParams params) { - alignas(128) extern __shared__ uint8_t smem[]; - __shared__ int32_t s_topk_indices[K]; - const auto batch_id = blockIdx.x; - const auto seq_len = params.seq_lens[batch_id]; - const auto transform = params.get_transform(batch_id, s_topk_indices); - // trivial case - if (seq_len <= K) { - impl::trivial_transform(transform, seq_len, K); - } else { - Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem, /*use_pdl=*/true); - device::PDLTriggerSecondary(); - Small::transform(transform); - } -} - -LARGE_TOPK_STAGE_1 void // long context, middle to large batch size -topk_combine_preprocess(const __grid_constant__ TopKParams params) { - alignas(128) extern __shared__ uint8_t smem[]; - __shared__ int32_t s_topk_indices[K]; - uint32_t work_id = blockIdx.x; - uint32_t batch_id; - uint32_t seq_len; - bool has_next; - uint32_t length; - uint32_t offset; - const auto cluster_rank = blockIdx.y; - - const auto prefetch_metadata = [&] { - const auto metadata = params.get_item_metadata(work_id); - batch_id = metadata.batch_id; - seq_len = metadata.seq_len; - has_next = metadata.has_next; - work_id += kNumClusters; // advance to the next item for this cluster - }; - const auto launch_prologue = [&] { - const auto partition = partition_work(seq_len, cluster_rank); - offset = partition.x; - length = partition.y; - Large::stage1_prologue(params.get_scores(batch_id) + offset, length, smem); - }; - - device::PDLWaitPrimary(); - device::PDLTriggerSecondary(); - - prefetch_metadata(); - if (seq_len == 0) return; - Large::stage1_init(smem); - launch_prologue(); - while (true) { - const auto this_length = length; - const auto this_offset = offset; - const auto need_prefetch = has_next; - const auto transform = params.get_transform(batch_id, s_topk_indices); - const auto ws = params.workspace + batch_id * params.workspace_stride; - if (need_prefetch) prefetch_metadata(); - Large::stage1(s_topk_indices, this_length, smem, /*reuse=*/true); - if (need_prefetch) launch_prologue(); - Large::stage1_epilogue(transform, this_offset, ws, smem); - if (!need_prefetch) break; - } -} - -LARGE_TOPK_STAGE_2 void // long context, middle to large batch size -topk_combine_transform(const __grid_constant__ TopKParams params) { - alignas(128) extern __shared__ uint8_t smem[]; - __shared__ int32_t s_topk_indices[K]; - const auto batch_id = blockIdx.x; - const auto seq_len = params.seq_lens[batch_id]; - const auto cluster_threshold = params.get_global_metadata().cluster_threshold; - const auto transform = params.get_transform(batch_id, s_topk_indices); - if (seq_len <= K) { - impl::trivial_transform(transform, seq_len, K); - } else if (seq_len <= kMax2PassLength) { - if (seq_len <= Small::kMax1PassLength) { - Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem); - } else { - __syncwarp(); - Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem); - } - Small::transform(transform); - } else if (seq_len <= cluster_threshold) { - Medium::run(params.get_scores(batch_id), seq_len, s_topk_indices, smem); - Medium::transform(transform, smem); - } else { - const auto ws = params.workspace + batch_id * params.workspace_stride; - device::PDLWaitPrimary(); - Large::transform(transform, ws, smem); - } -} - -FUSED_COMBINE_KERNEL void // long context, small batch size -topk_fused_transform(const __grid_constant__ TopKParams params) { - alignas(128) extern __shared__ uint8_t smem[]; - __shared__ int32_t s_topk_indices[K]; - const auto batch_id = blockIdx.x; - const auto cluster_rank = blockIdx.y; - const auto seq_len = params.seq_lens[batch_id]; - const auto transform = params.get_transform(batch_id, s_topk_indices); - if (seq_len <= K) { - if (cluster_rank != 0) return; // only first rank work - impl::trivial_transform(transform, seq_len, K); - } else if (seq_len <= Small::kMax1PassLength) { - if (cluster_rank != 0) return; // only first rank work - Small::run(params.get_scores(batch_id), s_topk_indices, seq_len, smem, /*use_pdl=*/true); - Small::transform(transform); - } else { - const auto [offset, length] = partition_work(seq_len, cluster_rank); - const auto ws = params.workspace + batch_id * params.workspace_stride; - Large::stage1_init(smem); - device::PDLWaitPrimary(); - Large::stage1_prologue(params.get_scores(batch_id) + offset, length, smem); - Large::stage1(s_topk_indices, length, smem); - Large::stage1_epilogue(transform, offset, ws, smem); - cooperative_groups::this_cluster().sync(); - if (cluster_rank != 0) return; // only first rank do the stage-2 - Large::transform(transform, ws, smem); - } -} - -struct CombinedTopKKernel { - static constexpr auto kStage1SMEM = sizeof(Large::Smem) + 128; - static constexpr auto kStage2SMEM = std::max(sizeof(Small::Smem), sizeof(Medium::Smem)) + 128; - - static void plan( // - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView metadata, - const uint32_t static_cluster_threshold) { - using namespace host; - auto B = SymbolicSize{"batch_size"}; - auto Bp1 = SymbolicSize{"batch_size_plus_1"}; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(seq_lens); - TensorMatcher({Bp1, 4}) // - .with_dtype() - .with_device(device_) - .verify(metadata); - - const auto batch_size = static_cast(B.unwrap()); - RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1); - if (batch_size <= kNumClusters) return; // metadata unused in fused path - - const auto device = device_.unwrap(); - constexpr auto kernel = topk_plan; - LaunchKernel(1, kBlockSize, device)( // - kernel, - static_cast(seq_lens.data_ptr()), - static_cast(metadata.data_ptr()), - batch_size, - static_cluster_threshold); - } - - static void transform( - const tvm::ffi::TensorView scores, - const tvm::ffi::TensorView seq_lens, - const tvm::ffi::TensorView page_table, - const tvm::ffi::TensorView page_indices, - const uint32_t page_size, - const tvm::ffi::TensorView workspace, - const tvm::ffi::TensorView metadata) { - using namespace host; - auto B = SymbolicSize{"batch_size"}; - auto Bp1 = SymbolicSize{"batch_size_plus_1"}; - auto L = SymbolicSize{"max_seq_len"}; - auto S = SymbolicSize{"score_stride"}; - auto P = SymbolicSize{"page_table_stride"}; - auto W = SymbolicSize{"workspace_stride"}; - constexpr auto D = Large::kWorkspaceInts; - auto device_ = SymbolicDevice{}; - device_.set_options(); - - TensorMatcher({B, L}) // - .with_strides({S, 1}) - .with_dtype() - .with_device(device_) - .verify(scores); - TensorMatcher({B}) // - .with_dtype() - .with_device(device_) - .verify(seq_lens); - TensorMatcher({B, -1}) // - .with_strides({P, 1}) - .with_dtype() - .with_device(device_) - .verify(page_table); - TensorMatcher({B, K}) // - .with_dtype() - .with_device(device_) - .verify(page_indices); - TensorMatcher({B, D}) // - .with_strides({W, 1}) - .with_dtype() - .with_device(device_) - .verify(workspace); - TensorMatcher({Bp1, 4}) // - .with_dtype() - .with_device(device_) - .verify(metadata); - - const auto page_bits = static_cast(std::countr_zero(page_size)); - const auto batch_size = static_cast(B.unwrap()); - const auto max_seq_len = static_cast(L.unwrap()); - const auto device = device_.unwrap(); - RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2"); - RuntimeCheck(S.unwrap() % 4 == 0, "score_stride must be a multiple of 4 (TMA 16-byte alignment)"); - RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1, "invalid metadata shape"); - - // NOTE: this should be fixed later - // RuntimeCheck(max_seq_len <= kMaxSupportedLength, max_seq_len, " exceeds the maximum supported length"); - - const auto params = TopKParams{ - .seq_lens = static_cast(seq_lens.data_ptr()), - .scores = static_cast(scores.data_ptr()), - .page_table = static_cast(page_table.data_ptr()), - .page_indices = static_cast(page_indices.data_ptr()), - .score_stride = S.unwrap(), - .page_table_stride = P.unwrap(), - .workspace = static_cast(workspace.data_ptr()), - .metadata = static_cast(metadata.data_ptr()), - .workspace_stride = W.unwrap() * static_cast(sizeof(int32_t)), - .batch_size = batch_size, - .page_bits = page_bits, - }; - - if (max_seq_len <= Small::kMax1PassLength) { - // All items fit in the short path -- no stage-1 needed - constexpr auto kernel = topk_short_transform; - setup_kernel_smem_once(); - LaunchKernel(batch_size, kBlockSize, device, kStage2SMEM) // - .enable_pdl(true)(kernel, params); - } else { - // Some items may be large -- launch stage-1 + main - if (batch_size <= kNumClusters) { - // can fuse into 1 stage - constexpr auto kernel = topk_fused_transform; - constexpr auto kSMEM = std::max(kStage1SMEM, kStage2SMEM); - setup_kernel_smem_once(); - LaunchKernel({batch_size, kClusterSize}, kBlockSize, device, kSMEM) - .enable_cluster({1, kClusterSize}) - .enable_pdl(true)(kernel, params); - } else { - // stage 1 + stage 2 - constexpr auto kernel_stage_1 = topk_combine_preprocess; - setup_kernel_smem_once(); - const auto num_clusters = std::min(batch_size, kNumClusters); - LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device, kStage1SMEM) - .enable_cluster({1, kClusterSize}) - .enable_pdl(true)(kernel_stage_1, params); - constexpr auto kernel_stage_2 = topk_combine_transform; - setup_kernel_smem_once(); - LaunchKernel(batch_size, kBlockSize, device, kStage2SMEM) // - .enable_pdl(true)(kernel_stage_2, params); - } - } - } -}; - -} // namespace diff --git a/lightllm/third_party/sglang_jit/dsv4/__init__.py b/lightllm/third_party/sglang_jit/dsv4/__init__.py deleted file mode 100644 index 507b225167..0000000000 --- a/lightllm/third_party/sglang_jit/dsv4/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -from .elementwise import fused_k_norm_rope_flashmla, fused_q_norm_rope -from .topk import topk_transform_512 - -__all__ = [ - "fused_k_norm_rope_flashmla", - "fused_q_norm_rope", - "topk_transform_512", -] diff --git a/lightllm/third_party/sglang_jit/dsv4/elementwise.py b/lightllm/third_party/sglang_jit/dsv4/elementwise.py deleted file mode 100644 index 07011b0479..0000000000 --- a/lightllm/third_party/sglang_jit/dsv4/elementwise.py +++ /dev/null @@ -1,215 +0,0 @@ -from typing import Optional, Tuple - -import torch - -from lightllm.third_party.sglang_jit.jit_utils import ( - cache_once, - is_arch_support_pdl, - load_jit, - make_cpp_args, -) -from lightllm.third_party.sglang_jit.runtime_utils import is_hip - -from .utils import make_name - -_is_hip = is_hip() - - -@cache_once -def _jit_fused_rope_module(): - args = make_cpp_args(is_arch_support_pdl()) - return load_jit( - make_name("fused_rope"), - *args, - cuda_files=["deepseek_v4/rope.cuh"], - cuda_wrappers=[("forward", f"FusedQKRopeKernel<{args}>::forward")], - ) - - -@cache_once -def _jit_main_q_norm_rope_module( - dtype: torch.dtype, - head_dim: int, - rope_dim: int, -): - """Main MLA path Q kernel: rmsnorm-self + RoPE, warp per (token, head).""" - args = make_cpp_args(dtype, head_dim, rope_dim, is_arch_support_pdl()) - return load_jit( - make_name("main_q_norm_rope"), - *args, - cuda_files=["deepseek_v4/main_norm_rope.cuh"], - cuda_wrappers=[ - ("forward", f"FusedQNormRopeKernel<{args}>::forward"), - ], - ) - - -@cache_once -def _jit_main_k_norm_rope_flashmla_module( - dtype: torch.dtype, - head_dim: int, - rope_dim: int, - page_size: int, -): - """Main MLA path K kernel: rmsnorm + RoPE + write to FlashMLA paged cache.""" - args = make_cpp_args(dtype, head_dim, rope_dim, page_size, is_arch_support_pdl()) - return load_jit( - make_name("main_k_norm_rope_flashmla"), - *args, - cuda_files=["deepseek_v4/main_norm_rope.cuh"], - cuda_wrappers=[ - ("forward", f"FusedKNormRopeFlashMLAKernel<{args}>::forward"), - ], - ) - - -@cache_once -def _jit_main_q_indexer_rope_hadamard_quant_module(dtype: torch.dtype): - """C4 indexer Q kernel: RoPE + 128-pt Hadamard + fp8 act-quant""" - args = make_cpp_args(dtype, is_arch_support_pdl()) - return load_jit( - make_name("main_q_indexer_rope_hadamard_quant"), - *args, - cuda_files=["deepseek_v4/main_norm_rope.cuh"], - cuda_wrappers=[ - ("forward", f"FusedQIndexerRopeHadamardQuantKernel<{args}>::forward"), - ], - ) - - -@cache_once -def _jit_main_q_indexer_rope_hadamard_fp4_quant_module(dtype: torch.dtype): - args = make_cpp_args(dtype, is_arch_support_pdl()) - return load_jit( - make_name("main_q_indexer_rope_hadamard_fp4_quant"), - *args, - cuda_files=["deepseek_v4/main_norm_rope.cuh"], - cuda_wrappers=[ - ("forward", f"FusedQIndexerRopeHadamardFp4QuantKernel<{args}>::forward"), - ], - ) - - -def fused_rope_inplace( - q: torch.Tensor, - k: Optional[torch.Tensor], - freqs_cis: torch.Tensor, - positions: torch.Tensor, - inverse: bool = False, -) -> None: - """Apply rotary embeddings to both Q and K in a single fused CUDA kernel. - - Args: - q: [batch_size, num_q_heads, rope_dim] bfloat16 - k: [batch_size, num_k_heads, rope_dim] bfloat16 or None - freqs_cis: [max_seq_len, rope_dim // 2] complex64 (full table) - positions: [batch_size] int32 or int64, indices into freqs_cis - inverse: if True, apply inverse rotation (conjugate freqs) - """ - if _is_hip: - from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton - - apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse) - if k is not None: - apply_rotary_emb_triton(k, freqs_cis, positions=positions, inverse=inverse) - return - - freqs_real = torch.view_as_real(freqs_cis).flatten(-2).contiguous() - module = _jit_fused_rope_module() - module.forward(q, k, freqs_real, positions, inverse) - - -def fused_q_norm_rope( - q_input: torch.Tensor, - q_output: torch.Tensor, - eps: float, - freqs_cis: torch.Tensor, - positions: torch.Tensor, -) -> None: - freqs_real = torch.view_as_real(freqs_cis).flatten(-2) - head_dim = q_input.shape[-1] - rope_dim = freqs_real.shape[-1] - module = _jit_main_q_norm_rope_module(q_input.dtype, head_dim, rope_dim) - module.forward(q_input, q_output, freqs_real, positions, eps) - - -def fused_q_indexer_rope_hadamard_quant( - q_input: torch.Tensor, - weight: torch.Tensor, - weight_scale: float, - freqs_cis: torch.Tensor, - positions: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - freqs_real = torch.view_as_real(freqs_cis).flatten(-2) - q_fp8 = torch.empty(q_input.shape, dtype=torch.float8_e4m3fn, device=q_input.device) - weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) - if _is_hip: - torch.ops.sgl_kernel.dsv4_fused_q_indexer_rope_hadamard_quant( - q_input, - q_fp8, - weight, - weights_out, - float(weight_scale), - freqs_real, - positions, - ) - else: - module = _jit_main_q_indexer_rope_hadamard_quant_module(q_input.dtype) - module.forward( - q_input, - q_fp8, - weight, - weights_out, - float(weight_scale), - freqs_real, - positions, - ) - return q_fp8, weights_out - - -def fused_q_indexer_rope_hadamard_fp4_quant( - q_input: torch.Tensor, - weight: torch.Tensor, - weight_scale: float, - freqs_cis: torch.Tensor, - positions: torch.Tensor, -) -> Tuple[Tuple[torch.Tensor, torch.Tensor], torch.Tensor]: - if _is_hip: - raise RuntimeError("DeepSeek V4 FP4 indexer requires the CUDA fused Q path.") - freqs_real = torch.view_as_real(freqs_cis).flatten(-2) - q_fp4 = torch.empty( - (*q_input.shape[:-1], q_input.shape[-1] // 2), - dtype=torch.int8, - device=q_input.device, - ) - q_sf = torch.empty(q_input.shape[:-1], dtype=torch.int32, device=q_input.device) - weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) - module = _jit_main_q_indexer_rope_hadamard_fp4_quant_module(q_input.dtype) - module.forward( - q_input, - q_fp4, - q_sf, - weight, - weights_out, - float(weight_scale), - freqs_real, - positions, - ) - return (q_fp4, q_sf), weights_out - - -def fused_k_norm_rope_flashmla( - kv: torch.Tensor, - kv_weight: torch.Tensor, - eps: float, - freqs_cis: torch.Tensor, - positions: torch.Tensor, - out_loc: torch.Tensor, - kvcache: torch.Tensor, - page_size: int, -) -> None: - freqs_real = torch.view_as_real(freqs_cis).flatten(-2) - head_dim = kv.shape[-1] - rope_dim = freqs_real.shape[-1] - module = _jit_main_k_norm_rope_flashmla_module(kv.dtype, head_dim, rope_dim, page_size) - module.forward(kv, kv_weight, freqs_real, positions, out_loc, kvcache, eps) diff --git a/lightllm/third_party/sglang_jit/dsv4/topk.py b/lightllm/third_party/sglang_jit/dsv4/topk.py deleted file mode 100644 index 1bfce7cef3..0000000000 --- a/lightllm/third_party/sglang_jit/dsv4/topk.py +++ /dev/null @@ -1,92 +0,0 @@ -from __future__ import annotations - -from typing import Optional - -import torch - -from lightllm.third_party.sglang_jit.jit_utils import ( - cache_once, - is_arch_support_pdl, - is_hip_runtime, - load_jit, - make_cpp_args, -) - -from .utils import make_name - - -@cache_once -def _jit_topk_v1_module(topk: int): - args = make_cpp_args(is_arch_support_pdl()) - assert topk in (512, 1024), "Only support topk=512 or 1024" - return load_jit( - make_name(f"topk_v1_{topk}"), - *args, - cuda_files=["deepseek_v4/topk_v1.cuh"], - cuda_wrappers=[("topk_transform", f"TopKKernel<{args}>::transform")], - extra_cuda_cflags=[f"-DSGL_TOPK={topk}"], - ) - - -@cache_once -def _jit_topk_v2_module(topk: int): - return load_jit( - make_name(f"topk_v2_{topk}"), - cuda_files=["deepseek_v4/topk_v2.cuh"], - cuda_wrappers=[ - ("topk_transform", "CombinedTopKKernel::transform"), - ("topk_plan", "CombinedTopKKernel::plan"), - ], - extra_cuda_cflags=[f"-DSGL_TOPK={topk}"], - ) - - -def topk_transform_512( - scores: torch.Tensor, - seq_lens: torch.Tensor, - page_tables: torch.Tensor, - out_page_indices: torch.Tensor, - page_size: int, - out_raw_indices: Optional[torch.Tensor] = None, -) -> None: - if is_hip_runtime(): - torch.ops.sgl_kernel.deepseek_v4_topk_transform_512( - scores, seq_lens, page_tables, out_page_indices, page_size, out_raw_indices - ) - else: - module = _jit_topk_v1_module(out_page_indices.shape[1]) - module.topk_transform(scores, seq_lens, page_tables, out_page_indices, page_size, out_raw_indices) - - -_WORKSPACE_INTS_PER_BATCH = 2 + 1024 * 2 -_PLAN_METADATA_INTS_PER_BATCH = 4 - - -def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor: - module = _jit_topk_v2_module(512) # does not matter - bs = seq_lens.shape[0] - metadata = seq_lens.new_empty(bs + 1, _PLAN_METADATA_INTS_PER_BATCH) - module.topk_plan(seq_lens, metadata, static_threshold) - return metadata - - -def topk_transform_512_v2( - scores: torch.Tensor, - seq_lens: torch.Tensor, - page_tables: torch.Tensor, - out_page_indices: torch.Tensor, - page_size: int, - metadata: torch.Tensor, -) -> None: - module = _jit_topk_v2_module(out_page_indices.shape[1]) - bs = scores.shape[0] - workspace = seq_lens.new_empty(bs, _WORKSPACE_INTS_PER_BATCH) - module.topk_transform( - scores, - seq_lens, - page_tables, - out_page_indices, - page_size, - workspace, - metadata, - ) diff --git a/lightllm/third_party/sglang_jit/dsv4/utils.py b/lightllm/third_party/sglang_jit/dsv4/utils.py deleted file mode 100644 index 8085074f6c..0000000000 --- a/lightllm/third_party/sglang_jit/dsv4/utils.py +++ /dev/null @@ -1,2 +0,0 @@ -def make_name(name: str) -> str: - return f"dpsk_v4_{name}" diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh deleted file mode 100644 index c9da765f4a..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/atomic.cuh +++ /dev/null @@ -1,35 +0,0 @@ -/// \file atomic.cuh -/// \brief Device-side atomic operations. - -#pragma once -#include - -namespace device::atomic { - -/** - * \brief Atomically computes the maximum of `*addr` and `value`, storing the - * result in `*addr`. - * \param addr Pointer to the value in global/shared memory to be updated. - * \param value The value to compare against. - * \return The old value at `*addr` before the update. - * \note On CUDA, this uses `atomicMax`/`atomicMin` on the reinterpreted - * integer representation. On ROCm, a CAS loop is used as a fallback. - */ -SGL_DEVICE float max(float* addr, float value) { -#ifndef USE_ROCM - float old; - old = (value >= 0) ? __int_as_float(atomicMax((int*)addr, __float_as_int(value))) - : __uint_as_float(atomicMin((unsigned int*)addr, __float_as_uint(value))); - return old; -#else - int* addr_as_i = (int*)addr; - int old = *addr_as_i, assumed; - do { - assumed = old; - old = atomicCAS(addr_as_i, assumed, __float_as_int(fmaxf(value, __int_as_float(assumed)))); - } while (assumed != old); - return __int_as_float(old); -#endif -} - -} // namespace device::atomic diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh deleted file mode 100644 index b47a4a27b2..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/cta.cuh +++ /dev/null @@ -1,40 +0,0 @@ -/// \file cta.cuh -/// \brief CTA (Cooperative Thread Array / thread-block) level primitives. - -#pragma once -#include -#include -#include - -namespace device::cta { - -/** - * \brief Compute the maximum of `value` across all threads in the CTA. - * - * Uses a two-level reduction: first within each warp via `warp::reduce_max`, - * then across warps using shared memory. The final result is stored in - * `smem[0]`. - * - * \tparam T Numeric type (must be supported by `warp::reduce_max`). - * \param value Per-thread input value. - * \param smem Shared memory buffer (must have at least `blockDim.x / 32` - * elements). - * \param min_value Identity element for max (default 0.0f). - * \note This function does NOT issue a trailing `__syncthreads()`. - * Callers must synchronize before reading `smem[0]`. - */ -template -SGL_DEVICE void reduce_max(T value, float* smem, float min_value = 0.0f) { - const uint32_t warp_id = threadIdx.x / kWarpThreads; - smem[warp_id] = warp::reduce_max(value); - __syncthreads(); - if (warp_id == 0) { - const auto tx = threadIdx.x; - const auto local_value = tx * kWarpThreads < blockDim.x ? smem[tx] : min_value; - const auto max_value = warp::reduce_max(local_value); - smem[0] = max_value; - } - // no extra sync; it is caller's responsibility to sync if needed -} - -} // namespace device::cta diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh deleted file mode 100644 index 02b166d01c..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress.cuh +++ /dev/null @@ -1,37 +0,0 @@ -#pragma once - -#include - -#include - -#include -#include - -#include - -namespace device::compress { - -struct alignas(16) PrefillPlan { - uint32_t ragged_id; - uint32_t batch_id; - uint32_t position; - uint32_t window_len; // must be in `[0, compress_ratio * (1 + is_overlap))` - - bool is_valid(const uint32_t ratio, const bool is_overlap) const { - const uint32_t max_window_len = ratio * (1 + is_overlap); - return window_len < max_window_len; - } -}; - -} // namespace device::compress - -namespace host::compress { - -using device::compress::PrefillPlan; -using PrefillPlanTensorDtype = uint8_t; -inline constexpr int64_t kPrefillPlanDim = 16; - -static_assert(alignof(PrefillPlan) == sizeof(PrefillPlan)); -static_assert(sizeof(PrefillPlan) == kPrefillPlanDim * sizeof(PrefillPlanTensorDtype)); - -} // namespace host::compress diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh deleted file mode 100644 index 3e87127c5f..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/compress_v2.cuh +++ /dev/null @@ -1,99 +0,0 @@ -#pragma once - -#include -#include - -#include - -#include -#include - -#include - -namespace device::compress { - -/// \brief Per-batch decode plan. Layout: 16 bytes. -struct alignas(16) DecodePlan { - uint32_t seq_len; - int32_t write_loc; - int32_t read_page_0; - int32_t read_page_1; -}; - -/// \brief Per-token compress plan (used by c4/c128 prefill). Layout: 16 bytes. -struct alignas(16) CompressPlan { - uint32_t seq_len; - uint16_t ragged_id; - uint16_t buffer_len; - int32_t read_page_0; - /// \brief Stage 0 (CPU): batch_id (used to look up page table). - /// \brief Stage 1 (GPU): final state-pool write location. - int32_t read_page_1; - - static SGL_DEVICE __host__ CompressPlan invalid() { - return CompressPlan{-1u, 0, 0, -1, -1}; - } - - SGL_DEVICE __host__ bool is_invalid() const { - return seq_len == -1u; - } -}; - -/// \brief Per-token write plan (used by c4/c128 prefill). Layout: 8 bytes. -struct alignas(8) WritePlan { - /// \brief Stage 0 (CPU): packed `(batch_id << 16) | ragged_id`. - /// \brief Stage 1 (GPU): just `ragged_id`. - uint32_t ragged_id; - /// \brief Stage 0 (CPU): position + 1 (used to look up state slot). - /// \brief Stage 1 (GPU): final state-pool write location. - int32_t write_loc; - - static SGL_DEVICE __host__ WritePlan invalid() { - return WritePlan{-1u, -1}; - } - - SGL_DEVICE __host__ bool is_invalid() const { - return ragged_id == -1u; - } -}; - -} // namespace device::compress - -namespace host::compress { - -using device::compress::CompressPlan; -using device::compress::DecodePlan; -using device::compress::WritePlan; - -static_assert(alignof(DecodePlan) == sizeof(DecodePlan)); -static_assert(sizeof(DecodePlan) == 16); -static_assert(alignof(CompressPlan) == sizeof(CompressPlan)); -static_assert(sizeof(CompressPlan) == 16); -static_assert(alignof(WritePlan) == sizeof(WritePlan)); -static_assert(sizeof(WritePlan) == 8); - -inline auto verify_plan_d(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const DecodePlan* { - TensorMatcher({N, sizeof(DecodePlan)}) // - .with_dtype() - .with_device(device) - .verify(t); - return static_cast(t.data_ptr()); -} - -inline auto verify_plan_c(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const CompressPlan* { - TensorMatcher({N, sizeof(CompressPlan)}) // - .with_dtype() - .with_device(device) - .verify(t); - return static_cast(t.data_ptr()); -} - -inline auto verify_plan_w(tvm::ffi::TensorView t, SymbolicSize& N, SymbolicDevice& device) -> const WritePlan* { - TensorMatcher({N, sizeof(WritePlan)}) // - .with_dtype() - .with_device(device) - .verify(t); - return static_cast(t.data_ptr()); -} - -} // namespace host::compress diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh deleted file mode 100644 index 53a62755b4..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/fp8_utils.cuh +++ /dev/null @@ -1,112 +0,0 @@ -#pragma once - -#include -#include -#include - -#include -#ifndef USE_ROCM -#include -#endif - -// Small helpers shared by the DeepSeek-V4 FP8/UE8M0 quantization kernels -// (silu_and_mul_masked_post_quant, store, mega_moe_pre_dispatch, ...). -// All functions are `SGL_DEVICE` (= `__forceinline__ __device__`) so -// including this header in multiple translation units is ODR-safe. - -namespace deepseek_v4::fp8 { - -// Round `x` to the nearest representable UE8M0 value. Returns the raw -// 8-bit biased exponent; the actual fp32 scale is `2^(exp - 127)` -// (i.e. `__uint_as_float(exp << 23)`). -SGL_DEVICE int32_t cast_to_ue8m0(float x) { - uint32_t u = __float_as_uint(x); - int32_t exp = int32_t((u >> 23) & 0xFF); - uint32_t mant = u & 0x7FFFFF; - return exp + (mant != 0); -} - -// 1 / 2^(exp - 127) as fp32. Equivalent to `1.0f / __uint_as_float(exp << 23)`. -SGL_DEVICE float inv_scale_ue8m0(int32_t exp) { - return __uint_as_float((127 + 127 - exp) << 23); -} - -// Clamp to [-FP8_E4M3_MAX, FP8_E4M3_MAX]. -// Uses platform-specific max from type.cuh (448 for E4M3FN, 224 for E4M3FNUZ). -SGL_DEVICE float fp8_e4m3_clip(float val) { - return fmaxf(fminf(val, kFP8E4M3Max), -kFP8E4M3Max); -} - -#ifndef USE_ROCM -// Pack two fp32 values into a single fp8x2_e4m3 with clamping. -SGL_DEVICE fp8x2_e4m3_t pack_fp8(float x, float y) { - return fp8x2_e4m3_t{fp32x2_t{fp8_e4m3_clip(x), fp8_e4m3_clip(y)}}; -} -#else -// Software float -> FP8 E4M3 conversion for ROCm/HIP. -// Supports both E4M3FN (MI350X, gfx950) and E4M3FNUZ (MI300X, gfx942). -SGL_DEVICE uint8_t cvt_float_to_fp8_e4m3(float val) { - val = fp8_e4m3_clip(val); - if (val == 0.0f) return 0; - - uint32_t f32 = __float_as_uint(val); - uint8_t sign = static_cast((f32 >> 31) << 7); - int32_t exp32 = static_cast((f32 >> 23) & 0xFF) - 127; - uint32_t mant23 = f32 & 0x7FFFFF; - -#if HIP_FP8_TYPE_FNUZ - // E4M3FNUZ: bias=8, max=240, no negative zero, NaN=0x80 - constexpr int32_t kBias = 8; - constexpr int32_t kMaxExp = 15; - constexpr int32_t kMinSubnormExp = -10; // min subnormal exponent - constexpr int32_t kMinNormExp = -7; // min normal exponent - constexpr uint8_t kSaturate = 0x7Fu; // max normal = 0_1111_111 = 240.0 -#else - // E4M3FN: bias=7, max=448, NaN=0x7F - constexpr int32_t kBias = 7; - constexpr int32_t kMaxExp = 15; - constexpr int32_t kMinSubnormExp = -9; - constexpr int32_t kMinNormExp = -6; - constexpr uint8_t kSaturate = 0x7Eu; // max normal = 0_1111_110 = 448.0 -#endif - - int32_t exp8; - uint8_t mant3; - - if (exp32 < kMinSubnormExp) { - return sign; - } else if (exp32 < kMinNormExp) { - // Subnormal range - int32_t shift = -(kBias - 1) - exp32; // 1..3 - uint32_t subnorm_mant = (0x800000 | mant23) >> (shift + 20); - uint32_t round_bit = ((0x800000 | mant23) >> (shift + 19)) & 1; - subnorm_mant += round_bit; - mant3 = static_cast(subnorm_mant & 0x07); - exp8 = 0; - if (subnorm_mant > 7) { - exp8 = 1; - mant3 = 0; - } - } else { - exp8 = exp32 + kBias; - mant3 = static_cast(mant23 >> 20); - uint32_t round_bit = (mant23 >> 19) & 1; - mant3 += round_bit; - if (mant3 > 7) { - mant3 = 0; - exp8++; - } - if (exp8 >= kMaxExp) return sign | kSaturate; - } - return sign | (static_cast(exp8) << 3) | mant3; -} - -// Pack two fp32 values into a single fp8x2_e4m3 (uint16_t on HIP). -SGL_DEVICE fp8x2_e4m3_t pack_fp8(float x, float y) { - uint8_t x8 = cvt_float_to_fp8_e4m3(x); - uint8_t y8 = cvt_float_to_fp8_e4m3(y); - return static_cast(x8) | (static_cast(y8) << 8); -} -#endif - -} // namespace deepseek_v4::fp8 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh deleted file mode 100644 index 0a3acc4773..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/kvcacheio.cuh +++ /dev/null @@ -1,96 +0,0 @@ -#include -#include - -#include - -#include - -namespace device::hisparse { - -/// NOTE: We call nope+rope as a "value" here. -/// GPU Cache layout: -/// VALUE 0, VALUE 1, ..., VALUE 63, -/// SCALE 0, SCALE 1, ..., SCALE 63, -/// [Padding to align to 576 bytes] -/// CPU Cache follow a trivial linear layout without any padding. -inline constexpr int64_t kGPUPageSize = 64; -inline constexpr int64_t kGPUPageBits = 6; // log2(kGPUPageSize) -inline constexpr int64_t kValueBytes = 576; -inline constexpr int64_t kScaleBytes = 8; -/// NOTE: FlashMLA requires each page to be aligned to 576 bytes -inline constexpr int64_t kCPUItemBytes = kValueBytes + kScaleBytes; -inline constexpr int64_t kGPUPageBytes = host::div_ceil(kCPUItemBytes * kGPUPageSize, 576) * 576; -inline constexpr int64_t kGPUScaleOffset = kValueBytes * kGPUPageSize; - -struct PointerInfo { - int64_t* value_ptr; - int64_t* scale_ptr; -}; - -SGL_DEVICE PointerInfo get_pointer_gpu(void* cache, int32_t index) { - using namespace device; - static_assert(1 << kGPUPageBits == kGPUPageSize); - const int32_t page_num = index >> kGPUPageBits; - const int32_t page_offset = index & (kGPUPageSize - 1); - const auto page_ptr = pointer::offset(cache, page_num * kGPUPageBytes); - const auto value_ptr = pointer::offset(page_ptr, page_offset * kValueBytes); - const auto scale_ptr = pointer::offset(page_ptr, kGPUScaleOffset + page_offset * kScaleBytes); - return {static_cast(value_ptr), static_cast(scale_ptr)}; -} - -SGL_DEVICE PointerInfo get_pointer_cpu(void* cache, int32_t index) { - using namespace device; - const auto value_ptr = pointer::offset(cache, index * kCPUItemBytes); - const auto scale_ptr = pointer::offset(value_ptr, kValueBytes); - return {static_cast(value_ptr), static_cast(scale_ptr)}; -} - -enum class TransferDirection { - DeviceToDevice = 0, - DeviceToHost = 1, - HostToDevice = 2, -}; - -template -SGL_DEVICE void transfer_item(void* dst_cache, void* src_cache, const int32_t dst_index, const int32_t src_index) { - constexpr bool is_dst_device = (direction != TransferDirection::DeviceToHost); - constexpr bool is_src_device = (direction != TransferDirection::HostToDevice); - constexpr auto dst_fn = is_dst_device ? get_pointer_gpu : get_pointer_cpu; - constexpr auto src_fn = is_src_device ? get_pointer_gpu : get_pointer_cpu; - - const auto [dst_value_ptr, dst_scale_ptr] = dst_fn(dst_cache, dst_index); - const auto [src_value_ptr, src_scale_ptr] = src_fn(src_cache, src_index); - - int64_t local_items[2]; - const int64_t* tail_src_ptr; - int64_t* tail_dst_ptr; - - const int32_t lane_id = threadIdx.x % 32; - - for (int i = 0; i < 2; ++i) { - const auto j = lane_id + i * 32; - local_items[i] = src_value_ptr[j]; - } - - if (lane_id < 8) { // handle the tail element safely - const auto last_id = 64 + lane_id; - tail_src_ptr = src_value_ptr + last_id; - tail_dst_ptr = dst_value_ptr + last_id; - } else { // broadcast load/store is safe - tail_src_ptr = src_scale_ptr; - tail_dst_ptr = dst_scale_ptr; - } - - const auto tail_item = *tail_src_ptr; - - // store first 512 bytes of value - for (int i = 0; i < 2; ++i) { - const auto j = lane_id + i * 32; - dst_value_ptr[j] = local_items[i]; - } - - // store the tail element - *tail_dst_ptr = tail_item; -} - -} // namespace device::hisparse diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh deleted file mode 100644 index e58214c951..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/cluster.cuh +++ /dev/null @@ -1,257 +0,0 @@ -#pragma once -#include -#include -#include - -#include "common.cuh" -#include "ptx.cuh" -#include -#include - -namespace device::top512 { - -template -struct ClusterTopK { - static constexpr uint32_t kClusterSize = 8; - static constexpr uint32_t kHistBits = 10; - static constexpr uint32_t kHistBins = 1 << kHistBits; - static constexpr uint32_t kRadixBins = 256; - static constexpr uint32_t kElemPerStage = 8; - static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize; - static constexpr uint32_t kNumStages = 4; - static constexpr uint32_t kMaxLength = kClusterSize * kNumStages * kSizePerStage; - static constexpr uint32_t kStoreLane = kBlockSize - 1; - static constexpr uint32_t kAboveBits = 11; - - // --------------------------------------------------------------------------- - // Shared memory layouts - // --------------------------------------------------------------------------- - - struct Smem { - uint64_t barrier[kNumStages]; - uint32_t local_above_equal[kClusterSize]; - uint32_t prefix_above_equal; - alignas(128) uint32_t counter_gt; - alignas(128) uint32_t counter_eq; - alignas(128) MatchBin match; - alignas(128) uint32_t warp_sum[kNumWarps]; - uint32_t histogram[kHistBins]; - alignas(128) float score_buffer[kNumStages][kSizePerStage]; - Tie tie_buffer[kMaxTies]; - }; - - struct alignas(16) Metadata { - uint32_t batch_id; - uint32_t seq_len; - bool has_next; - }; - - struct WorkSpace { - uint2 metadata; // {num_above, num_ties} - Tie ties[kMaxTies]; - }; - - static constexpr uint32_t kWorkspaceInts = sizeof(WorkSpace) / sizeof(uint32_t); - - // --------------------------------------------------------------------------- - // Stage 1: histogram + cluster reduce + find threshold + scatter - // --------------------------------------------------------------------------- - - SGL_DEVICE static void stage1_init(void* _smem) { - const auto tx = threadIdx.x; - __builtin_assume(tx < kBlockSize); - const auto smem = static_cast(_smem); - if (tx < kHistBins) smem->histogram[tx] = 0; - if (tx < kNumStages) ptx::mbarrier_init(&smem->barrier[tx], 1); - __syncthreads(); - } - - SGL_DEVICE static void stage1_prologue(const float* scores, uint32_t length, void* _smem) { - if (threadIdx.x == 0) { - const auto smem = static_cast(_smem); - const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage; - const auto length_aligned = (length + 3u) & ~3u; // align to 4 for TMA -#pragma unroll - for (uint32_t stage = 0; stage < kNumStages; stage++) { - if (stage >= num_stages) break; - const auto offset = stage * kSizePerStage; - const auto size = min(kSizePerStage, length_aligned - offset); - const auto size_bytes = size * sizeof(float); - const auto bar = &smem->barrier[stage]; - ptx::tma_load(smem->score_buffer[stage], scores + offset, size_bytes, bar); - ptx::mbarrier_arrive_expect_tx(bar, size_bytes); - } - } - } - - SGL_DEVICE static void stage1(int32_t* indices, uint32_t length, void* _smem, bool reuse = false) { - const auto smem = static_cast(_smem); - const auto tx = threadIdx.x; - __builtin_assume(tx < kBlockSize); - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - // Initialize shared memory histogram, counters, and barriers -#pragma unroll - for (uint32_t stage = 0; stage < kNumStages; stage++) { - const auto offset = stage * kSizePerStage; - if (offset >= length) break; - const auto size = min(kSizePerStage, length - offset); - if (lane_id == 0) ptx::mbarrier_wait(&smem->barrier[stage], 0); - __syncwarp(); -#pragma unroll - for (uint32_t i = 0; i < kElemPerStage; ++i) { - const auto idx = tx + i * kBlockSize; - if (idx >= size) break; - const auto score = smem->score_buffer[stage][idx]; - const auto bin = extract_coarse_bin(score); - atomicAdd(&smem->histogram[bin], 1); - } - } - - static_assert(kHistBins <= kBlockSize); - - // 2-shot all-reduce - { - auto cluster = cooperative_groups::this_cluster(); - cluster.sync(); - const auto cluster_rank = blockIdx.y; - const auto kLocalSize = kHistBins / kClusterSize; - const auto offset = kLocalSize * cluster_rank; - - const auto src_tx = tx / kClusterSize; - const auto src_rank = tx % kClusterSize; - - if (tx < kHistBins) { - const auto addr = &smem->histogram[offset + src_tx]; - const auto src_addr = cluster.map_shared_rank(addr, src_rank); - *src_addr = warp::reduce_sum(*src_addr); - } - cluster.sync(); - } - - // now each block holds the whole histogram, find the threshold bin - { - const auto value = tx < kHistBins ? smem->histogram[tx] : 0; - const auto warp_inc = warp_inclusive_sum(lane_id, value); - if (lane_id == kWarpThreads - 1) { - smem->warp_sum[warp_id] = warp_inc; - } - - __syncthreads(); - const auto tmp = smem->warp_sum[lane_id]; - // total_length = sum of all bins in the globally-reduced histogram - // (problem.length is block-local; after cluster reduction we need the global total) - const auto total_length = warp::reduce_sum(tmp); - uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); - prefix_sum += warp_inc; - const auto above = total_length - prefix_sum; - if (tx < kHistBins && above < K && above + value >= K) { - smem->counter_gt = smem->counter_eq = 0; - smem->match = { - .bin = tx, - .above_count = above, - .equal_count = value, - }; - } - __syncthreads(); - } - - const auto [thr_bin, num_above, num_equal] = smem->match; - - // write above and equal results to global memory -#pragma unroll - for (uint32_t stage = 0; stage < kNumStages; stage++) { - const auto offset = stage * kSizePerStage; - if (offset >= length) break; -#pragma unroll - for (uint32_t i = 0; i < kElemPerStage; ++i) { - const auto buf_idx = tx + i * kBlockSize; - const auto global_idx = offset + buf_idx; - if (global_idx >= length) break; - const auto score = smem->score_buffer[stage][buf_idx]; - const auto bin = extract_coarse_bin(score); - if (bin > thr_bin) { - indices[atomicAdd(&smem->counter_gt, 1)] = global_idx; - } else if (bin == thr_bin) { - const auto pos = atomicAdd(&smem->counter_eq, 1); - if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score}; - } - } - } - if (reuse) { - const auto num_stages = (length + kSizePerStage - 1) / kSizePerStage; - if (tx < kHistBins) smem->histogram[tx] = 0; - if (tx < num_stages) ptx::mbarrier_arrive(&smem->barrier[tx]); - } - __syncthreads(); - } - - // --------------------------------------------------------------------------- - // Stage 1 epilogue: cross-block prefix sum + page translate + tie store - // --------------------------------------------------------------------------- - - SGL_DEVICE static void stage1_epilogue(const TransformParams params, const uint32_t offset, void* _ws, void* _smem) { - auto cluster = cooperative_groups::this_cluster(); - const auto smem = static_cast(_smem); - const auto tx = threadIdx.x; - const auto local_above = smem->counter_gt; - const auto local_equal = smem->counter_eq; - const auto cluster_rank = blockIdx.y; - - constexpr uint32_t kAboveMask = (1 << kAboveBits) - 1; - static_assert(kAboveMask >= K); - - // Pack local counts -- NO alignment rounding (contiguous layout) - static_assert(kMaxTies <= kBlockSize); - const auto idx_above = tx < local_above ? params.indices_in[tx] : 0; - const auto tie_value = tx < local_equal ? smem->tie_buffer[tx] : Tie{0, 0.0f}; - - // push to remote shared memory, can reduce latency of reading remote - if (tx < kClusterSize) { - const auto value = (local_equal << kAboveBits) | local_above; - const auto dst_addr = cluster.map_shared_rank(smem->local_above_equal, tx); - dst_addr[cluster_rank] = value; - } - // after this last sync, only read local shared memory - // so that it is safe when peer rank has already exited the kernel - cluster.sync(); - if (tx < kClusterSize) { - const auto value = tx < cluster_rank ? smem->local_above_equal[tx] : 0; - const auto kActiveMask = (1u << kClusterSize) - 1; - smem->prefix_above_equal = warp::reduce_sum(value, kActiveMask); - } - __syncthreads(); - - const auto prefix_packed = smem->prefix_above_equal; - const auto prefix_above = prefix_packed & kAboveMask; - const auto prefix_equal = prefix_packed >> kAboveBits; - - // Page-translate above elements - if (tx < local_above) { - params.write(tx + prefix_above, idx_above + offset); - } - // Contiguous tie store via regular global writes (no TMA, no gaps) - const auto ws = static_cast(_ws); - if (tx < local_equal && tx + prefix_equal < kMaxTies) { - ws->ties[tx + prefix_equal] = {tie_value.idx + offset, tie_value.score}; - } - // Block 0 writes global metadata {num_above, num_ties} - if (cluster_rank == kClusterSize - 1 && tx == 0) { - const auto sum_above = prefix_above + local_above; - const auto sum_equal = prefix_equal + local_equal; - ws->metadata = make_uint2(sum_above, sum_equal); - } - } - - SGL_DEVICE static void transform(const TransformParams params, const void* _ws, void* _smem) { - const auto ws = static_cast(_ws); - const auto meta = &ws->metadata; - const auto [num_above, num_equal] = *meta; - if (num_above >= K || num_equal == 0) return; - const auto clamped_ties = min(num_equal, kMaxTies); - tie_handle_transform(ws->ties, clamped_ties, num_above, K, params, _smem); - } -}; - -} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh deleted file mode 100644 index d553032d79..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/common.cuh +++ /dev/null @@ -1,176 +0,0 @@ -#pragma once -#include -#include -#include -#include - -#include - -namespace device::top512 { - -inline constexpr uint32_t kMaxTopK = 1024; -inline constexpr uint32_t kBlockSize = 1024; -inline constexpr uint32_t kNumWarps = kBlockSize / kWarpThreads; -inline constexpr uint32_t kMaxTies = 1024; // == kBlockSize: 1 element per thread in stage2 -static constexpr uint32_t kRadixBins = 256; -static_assert(kMaxTopK <= kBlockSize && kMaxTies <= kBlockSize); - -// always use float4 to load from global memory -using Vec4 = AlignedVector; - -SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) { - const uint32_t mask = (1u << page_bits) - 1u; - return (page_table[i >> page_bits] << page_bits) | (i & mask); -} - -struct TransformParams { - const int32_t* __restrict__ page_table; - const int32_t* __restrict__ indices_in; - int32_t* __restrict__ indices_out; - uint32_t page_bits; - - SGL_DEVICE void transform(const uint32_t idx) const { - indices_out[idx] = page_to_indices(page_table, indices_in[idx], page_bits); - } - SGL_DEVICE void write(const uint32_t dst, const uint32_t src) const { - indices_out[dst] = page_to_indices(page_table, src, page_bits); - } -}; - -struct alignas(16) MatchBin { - uint32_t bin; - uint32_t above_count; - uint32_t equal_count; -}; - -struct alignas(8) Tie { - uint32_t idx; - float score; -}; - -struct TieHandleSmem { - alignas(128) uint32_t counter; // output position counter - alignas(128) MatchBin match; - uint32_t histogram[kRadixBins]; // 256-bin radix histogram - uint32_t warp_sum[kNumWarps]; // for 2-pass prefix sum -}; - -template -SGL_DEVICE uint32_t extract_coarse_bin(float x) { - static_assert(0 < kBits && kBits < 15); - const auto hx = cast(x); - const uint16_t bits = *reinterpret_cast(&hx); - const uint16_t key = (bits & 0x8000) ? ~bits : bits | 0x8000; - return key >> (16 - kBits); -} - -SGL_DEVICE uint32_t warp_inclusive_sum(uint32_t lane_id, uint32_t val) { - static_assert(kWarpThreads == 32); -#pragma unroll - for (uint32_t offset = 1; offset < 32; offset *= 2) { - uint32_t n = __shfl_up_sync(0xFFFFFFFF, val, offset); - if (lane_id >= offset) val += n; - } - return val; -} - -/// Order-preserving float32 -> uint32 for radix select -SGL_DEVICE uint32_t extract_exact_bin(float x) { - uint32_t bits = __float_as_uint(x); - return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u); -} - -SGL_DEVICE void trivial_transform(const TransformParams& params, uint32_t length, uint32_t K) { - if (const auto tx = threadIdx.x; tx < length) { - params.write(tx, tx); - } else if (tx < K) { - params.indices_out[tx] = -1; - } -} - -SGL_DEVICE void tie_handle_transform( - const Tie* __restrict__ ties, // - const uint32_t num_ties, - const uint32_t num_above, - const uint32_t K, - const TransformParams params, - void* _smem) { - auto* smem = static_cast(_smem); - const auto tx = threadIdx.x; - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - // Each thread loads one element (or becomes inactive) - const bool has_elem = tx < num_ties; - const auto tie = has_elem ? ties[tx] : Tie{0, 0.0f}; - const uint32_t key = extract_exact_bin(tie.score); - const uint32_t idx = tie.idx; - bool active = has_elem; - uint32_t topk_remain = K - num_above; - uint32_t write_pos = K; - - smem->counter = 0; - __syncthreads(); - - // Number of warps covering the 256-bin histogram (256/32 = 8) - constexpr uint32_t kRadixWarps = kRadixBins / kWarpThreads; - -#pragma unroll - for (int round = 0; round < 4; round++) { - const uint32_t shift = 24 - round * 8; - const uint32_t bin = (key >> shift) & 0xFFu; - - // 1. Build histogram - if (tx < kRadixBins) smem->histogram[tx] = 0; - __syncthreads(); - if (active) atomicAdd(&smem->histogram[bin], 1); - __syncthreads(); - - // 2. v2-style 2-pass prefix sum on 256 bins - // Only first 256 threads (8 warps) carry histogram bins. - // Other threads get hist_val=0 and harmless prefix results. - uint32_t hist_val = 0; - uint32_t warp_inc = 0; - if (tx < kRadixBins) { - hist_val = smem->histogram[tx]; - warp_inc = warp_inclusive_sum(lane_id, hist_val); - if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; - } - __syncthreads(); - if (tx < kRadixBins) { - // Inter-warp prefix (only first kHistWarps warp totals matter) - const auto tmp = (lane_id < kRadixWarps) ? smem->warp_sum[lane_id] : 0; - const auto total = warp::reduce_sum(tmp); - const auto inter = warp::reduce_sum(lane_id < warp_id ? tmp : 0); - const auto prefix = inter + warp_inc; // inclusive prefix through this bin - const auto above = total - prefix; // elements in bins ABOVE this one - // 3. Find threshold bin - if (above < topk_remain && above + hist_val >= topk_remain) { - smem->match = {tx, above, topk_remain - above}; - } - } - __syncthreads(); - - const auto [thr, n_above, _] = smem->match; - - // 4. Scatter - if (active) { - if (bin > thr) { - write_pos = num_above + atomicAdd(&smem->counter, 1); - active = false; - } else if (bin < thr) { - active = false; - } else if (round == 3) { - write_pos = K - atomicAdd(&smem->match.equal_count, -1u); - } - // my_bin == thr && round < 3: stay active for next round - } - - topk_remain -= n_above; - if (topk_remain == 0) break; - } - - if (write_pos < K) params.write(write_pos, idx); -} - -} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh deleted file mode 100644 index 73eef555f4..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/ptx.cuh +++ /dev/null @@ -1,54 +0,0 @@ -#pragma once -#include - -#include - -#include - -namespace device::top512 { - -namespace ptx { - -SGL_DEVICE void mbarrier_wait(uint64_t* addr, uint32_t phase) { - while (!cuda::ptx::mbarrier_try_wait_parity(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, addr, phase)) - ; -} - -SGL_DEVICE void mbarrier_init(uint64_t* addr, uint32_t arrives) { - cuda::ptx::mbarrier_init(addr, arrives); -} - -SGL_DEVICE void mbarrier_arrive_expect_tx(uint64_t* addr, uint32_t tx) { - cuda::ptx::mbarrier_arrive_expect_tx(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, cuda::ptx::space_shared, addr, tx); -} - -SGL_DEVICE void mbarrier_arrive(uint64_t* addr) { - cuda::ptx::mbarrier_arrive(cuda::ptx::sem_relaxed, cuda::ptx::scope_cta, cuda::ptx::space_shared, addr); -} - -SGL_DEVICE void tma_load(void* dst, const void* src, uint32_t num_bytes, uint64_t* mbar) { - cuda::ptx::cp_async_bulk(cuda::ptx::space_shared, cuda::ptx::space_global, dst, src, num_bytes, mbar); -} - -SGL_DEVICE uint32_t elect_sync() { - uint32_t pred = 0; - asm volatile( - "{\n\t" - ".reg .pred %%px;\n\t" - "elect.sync _|%%px, %1;\n\t" - "@%%px mov.s32 %0, 1;\n\t" - "}" - : "+r"(pred) - : "r"(0xFFFFFFFF)); - return pred; -} - -SGL_DEVICE bool elect_sync_cta(uint32_t tx) { - const auto warp_id = tx / 32; - const auto uniform_warp_id = __shfl_sync(0xFFFFFFFF, warp_id, 0); - return (uniform_warp_id == 0 && elect_sync()); -} - -} // namespace ptx - -} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh deleted file mode 100644 index 77d7361ee8..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/register.cuh +++ /dev/null @@ -1,302 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "common.cuh" -#include "ptx.cuh" -#include -#include - -namespace device::top512 { - -template -struct RegisterTopK { - static constexpr uint32_t kHistBits = 12; - static constexpr uint32_t kHistBins = 1 << kHistBits; - static constexpr uint32_t kVecsPerThread = 4; - static constexpr uint32_t kMaxTolerance = 0; - static constexpr uint32_t kMax1PassLength = kVecsPerThread * 4 * kBlockSize; - static constexpr uint32_t kMaxExtraLength = kMax1PassLength; - static constexpr uint32_t kMax2PassLength = kMax1PassLength + kMaxExtraLength; - - struct Smem { - using HistVec = AlignedVector; - alignas(128) uint32_t counter_gt; - alignas(128) uint32_t counter_eq; - uint64_t mbarrier; // for cp.async - MatchBin match; - uint32_t warp_sum[kNumWarps]; - union { - uint32_t histogram[kHistBins]; - HistVec histogram_vec[kBlockSize]; - Tie tie_buffer[kMaxTies]; - }; - alignas(16) float score_buffer[kMaxExtraLength]; - }; - - template - SGL_DEVICE static void - run(const float* scores, // - int32_t* indices, - const uint32_t length, - void* _smem, - const bool use_pdl = false) { - const auto smem = static_cast(_smem); - const auto tx = threadIdx.x; - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - // Initialize shared memory histogram - { - typename Smem::HistVec hist_vec; - hist_vec.fill(0); - smem->histogram_vec[tx] = hist_vec; - if (tx == 0) { - smem->counter_gt = smem->counter_eq = 0; - if constexpr (kIs2Pass) { - ptx::mbarrier_init(&smem->mbarrier, 1); - } - } - __syncthreads(); - } - - if (use_pdl) device::PDLWaitPrimary(); - - // Load scores into registers - Vec4 local[kVecsPerThread]; -#pragma unroll - for (uint32_t v = 0; v < kVecsPerThread; ++v) { - const uint32_t base = (tx + v * kBlockSize) * 4; - if (base >= length) break; - local[v].load(scores, tx + v * kBlockSize); - } - - // Fetch the next chunk of scores - if constexpr (kIs2Pass) { - if (ptx::elect_sync_cta(tx)) { - const auto length_aligned = (length + 3u - kMax1PassLength) & ~3u; - const auto size_bytes = length_aligned * sizeof(float); - ptx::tma_load(smem->score_buffer, scores + kMax1PassLength, size_bytes, &smem->mbarrier); - ptx::mbarrier_arrive_expect_tx(&smem->mbarrier, size_bytes); - } - __syncwarp(); // avoid warp divergence on - } - - // Accumulate histogram via shared-memory atomics -#pragma unroll - for (uint32_t v = 0; v < kVecsPerThread; ++v) { -#pragma unroll - for (uint32_t e = 0; e < 4; ++e) { - if constexpr (!kIs2Pass) { - const uint32_t idx = (tx + v * kBlockSize) * 4 + e; - if (idx >= length) goto LABEL_ACC_FINISH; - } - atomicAdd(&smem->histogram[extract_coarse_bin(local[v][e])], 1); - } - } - if constexpr (kIs2Pass) { - // 16K ~ 32K. `i` is a float4 index - if (lane_id == 0) ptx::mbarrier_wait(&smem->mbarrier, 0); - __syncwarp(); - for (uint32_t i = tx; i + kMax1PassLength < length; i += kBlockSize) { - const auto val = smem->score_buffer[i]; - atomicAdd(&smem->histogram[extract_coarse_bin(val)], 1); - } - } - [[maybe_unused]] LABEL_ACC_FINISH: - __syncthreads(); - - // Phase 2: Exclusive prefix scan -> find threshold bin - { - constexpr uint32_t kItems = kHistBins / kBlockSize; - uint32_t orig[kItems]; - const auto hist_vec = smem->histogram_vec[tx]; - uint32_t tmp_local_sum = 0; - -#pragma unroll - for (uint32_t i = 0; i < kItems; ++i) { - orig[i] = hist_vec[i]; - tmp_local_sum += orig[i]; - } - - const auto warp_inc = warp_inclusive_sum(lane_id, tmp_local_sum); - const auto warp_exc = warp_inc - tmp_local_sum; - if (lane_id == kWarpThreads - 1) { - smem->warp_sum[warp_id] = warp_inc; - } - - __syncthreads(); - - const auto tmp = smem->warp_sum[lane_id]; - // Exactly one bin satisfies: above < K && above + count >= K - uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); - prefix_sum += warp_exc; -#pragma unroll - for (uint32_t i = 0; i < kItems; ++i) { - prefix_sum += orig[i]; - const auto above = length - prefix_sum; - if (above < K && above + orig[i] >= K) { - smem->match = { - .bin = tx * kItems + i, - .above_count = above, - .equal_count = orig[i], - }; - } - } - __syncthreads(); - } - - const auto [thr_bin, num_above, num_equal] = smem->match; - - // Phase 3: Scatter - // Elements strictly above threshold go directly to output. - // Tied elements: simple path admits first-come; tiebreak path collects into tie_buffer. - const bool need_tiebreak = (num_equal + num_above > K + kMaxTolerance); - const auto topk_indices = indices; - const auto tie_buffer = smem->tie_buffer; - -#pragma unroll - for (uint32_t v = 0; v < kVecsPerThread; ++v) { -#pragma unroll - for (uint32_t e = 0; e < 4; ++e) { - const uint32_t idx = (tx + v * kBlockSize) * 4 + e; - if constexpr (!kIs2Pass) { - if (idx >= length) goto LABEL_SCATTER_DONE; - } - const uint32_t bin = extract_coarse_bin(local[v][e]); - if (bin > thr_bin) { - topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx; - } else if (bin == thr_bin) { - const auto pos = atomicAdd(&smem->counter_eq, 1); - if (need_tiebreak) { - if (pos < kMaxTies) { - tie_buffer[pos] = {.idx = idx, .score = local[v][e]}; - } - } else { - if (const auto which = pos + num_above; which < K) { - topk_indices[which] = idx; - } - } - } - } - // prefetch the next scores - if constexpr (kIs2Pass) { - local[v].load(smem->score_buffer, tx + v * kBlockSize); - } - } - - // 16K ~ 32K, already in registers: similar loop as above but read from smem->score_buffer - if constexpr (kIs2Pass) { -#pragma unroll - for (uint32_t v = 0; v < kVecsPerThread; ++v) { -#pragma unroll - for (uint32_t e = 0; e < 4; ++e) { - const uint32_t idx = (tx + v * kBlockSize) * 4 + e + kMax1PassLength; - if (idx >= length) goto LABEL_SCATTER_DONE; - const uint32_t bin = extract_coarse_bin(local[v][e]); - if (bin > thr_bin) { - topk_indices[atomicAdd(&smem->counter_gt, 1)] = idx; - } else if (bin == thr_bin) { - const auto pos = atomicAdd(&smem->counter_eq, 1); - if (need_tiebreak) { - if (pos < kMaxTies) { - tie_buffer[pos] = {.idx = idx, .score = local[v][e]}; - } - } else { - if (const auto which = pos + num_above; which < K) { - topk_indices[which] = idx; - } - } - } - } - } - } - - [[maybe_unused]] LABEL_SCATTER_DONE: - if (!need_tiebreak) return; - - // Phase 4: Tie-breaking within the threshold bin. - // Assume num_ties <= kBlockSize (at most 1 block of ties). - // Each thread takes one tied element, computes its rank (number of - // elements with strictly higher score, breaking exact float ties by - // original index), and writes to output if rank < topk_remain. - __syncthreads(); - static_assert(kMaxTies <= kBlockSize); - - const uint32_t num_ties = min(num_equal, kMaxTies); - const uint32_t topk_remain = K - num_above; - - const auto is_greater = [](const Tie& a, const Tie& b) { - return (a.score > b.score) || (a.score == b.score && a.idx < b.idx); - }; - - if (num_ties <= kWarpThreads) { - static_assert(kWarpThreads <= kNumWarps); - if (lane_id >= num_ties || warp_id >= num_ties) return; // some threads are idle - /// NOTE: use long long to avoid mask overflow when num_ties == 32 - const uint32_t mask = (1ull << num_ties) - 1u; - const auto tie = tie_buffer[lane_id]; - const auto target_tie = tie_buffer[warp_id]; - const bool pred = is_greater(tie, target_tie); - const auto rank = static_cast(__popc(__ballot_sync(mask, pred))); - if (lane_id == 0 && rank < topk_remain) { - topk_indices[num_above + rank] = target_tie.idx; - } - } else if (num_ties <= kWarpThreads * 2) { - // 64 x 64 topk implementation: each thread takes 2 elements - const auto lane_id_1 = lane_id + kWarpThreads; - const auto warp_id_1 = warp_id + kWarpThreads; - const auto invalid = Tie{.idx = 0xFFFFFFFF, .score = -FLT_MAX}; - const auto tie_0 = tie_buffer[lane_id]; - const auto tie_1 = lane_id_1 < num_ties ? tie_buffer[lane_id_1] : invalid; - if (true) { - const auto target = tie_buffer[warp_id]; - const bool pred_0 = is_greater(tie_0, target); - const bool pred_1 = is_greater(tie_1, target); - const auto rank_0 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_0))); - const auto rank_1 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_1))); - const auto rank = rank_0 + rank_1; - if (lane_id == 0 && rank < topk_remain) { - topk_indices[num_above + rank] = target.idx; - } - } - if (warp_id_1 < num_ties) { - const auto target = tie_buffer[warp_id_1]; - const bool pred_0 = is_greater(tie_0, target); - const bool pred_1 = is_greater(tie_1, target); - const auto rank_0 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_0))); - const auto rank_1 = static_cast(__popc(__ballot_sync(0xFFFFFFFF, pred_1))); - const auto rank = rank_0 + rank_1; - if (lane_id == 0 && rank < topk_remain) { - topk_indices[num_above + rank] = target.idx; - } - } - } else { - /// NOTE: Based on my observation, this path is very rarely reached - [[unlikely]]; - // Block-level: each thread reads from tie_buffer in shared memory - for (auto i = warp_id; i < num_ties; i += kNumWarps) { - const auto target_tie = tie_buffer[i]; - uint32_t local_rank = 0; - for (auto j = lane_id; j < num_ties; j += kWarpThreads) { - const auto tie = tie_buffer[j]; - if (is_greater(tie, target_tie)) local_rank++; - } - // sum the rank across the warp - const auto rank = warp::reduce_sum(local_rank); - if (lane_id == 0 && rank < topk_remain) { - topk_indices[num_above + rank] = target_tie.idx; - } - } - } - } - - SGL_DEVICE static void transform(const TransformParams params) { - __syncthreads(); - if (const auto tx = threadIdx.x; tx < K) params.transform(tx); - } -}; - -} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh deleted file mode 100644 index 4462b89a19..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/deepseek_v4/topk/streaming.cuh +++ /dev/null @@ -1,213 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "common.cuh" -#include "ptx.cuh" -#include -#include - -namespace device::top512 { - -template -struct StreamingTopK { - static constexpr uint32_t kHistBits = 12; - static constexpr uint32_t kHistBins = 1 << kHistBits; - static constexpr uint32_t kRadixBins = 256; - static constexpr uint32_t kElemPerStage = 8; - static constexpr uint32_t kSizePerStage = kElemPerStage * kBlockSize; - static constexpr uint32_t kNumStages = 2; // double buffer - - static constexpr uint32_t kHistItems = kHistBins / kBlockSize; // 4 - static_assert(kHistItems * kBlockSize == kHistBins); - using HistVec = AlignedVector; - - struct Smem { - uint64_t barrier[2][kNumStages]; - alignas(128) uint32_t counter_gt; - alignas(128) uint32_t counter_eq; - alignas(128) MatchBin match; - alignas(128) uint32_t warp_sum[kNumWarps]; - union { - uint32_t histogram[kHistBins]; - HistVec histogram_vec[kBlockSize]; - Tie tie_buffer[kMaxTies]; - }; - union { - float score_buffer[kNumStages][kSizePerStage]; - TieHandleSmem stage2; // reuse smem for tie handling in phase D - }; - }; - - // --------------------------------------------------------------------------- - // Helpers - // --------------------------------------------------------------------------- - - /// NOTE: length must be 4-aligned since we load 4 floats/thread. Caller should round up. - template - SGL_DEVICE static void issue_tma(const float* scores, uint32_t stage, uint32_t length, Smem* smem) { - const auto buf_idx = stage % kNumStages; - const auto offset = stage * kSizePerStage; - const auto size = min(kSizePerStage, length - offset); - const auto size_bytes = size * sizeof(float); - const auto bar = &smem->barrier[kIsScatter][buf_idx]; - ptx::tma_load(smem->score_buffer[buf_idx], scores + offset, size_bytes, bar); - ptx::mbarrier_arrive_expect_tx(bar, size_bytes); - } - - // --------------------------------------------------------------------------- - // Unified streaming pass. Used for both phase A (kIsScatter=false) and - // phase C (kIsScatter=true). Each buffer is reused across iterations via the - // reuse-arrive trick (same pattern as ClusterTopKImpl::stage1). - // --------------------------------------------------------------------------- - - template - SGL_DEVICE static void stream_pass( - const float* scores, - const uint32_t length, - const uint32_t thr_bin, // ignored when !kIsScatter - int32_t* s_topk_indices, // ignored when !kIsScatter - Smem* smem) { - const auto tx = threadIdx.x; - const auto num_iters = (length + kSizePerStage - 1) / kSizePerStage; - const auto lane_id = tx % kWarpThreads; - - // Initial double-buffer TMA prologue. - const auto length_aligned = (length + 3u) & ~3u; - if (tx == 0) { -#pragma unroll - for (uint32_t i = 0; i < kNumStages; i++) { - if (i >= num_iters) break; - issue_tma(scores, i, length_aligned, smem); - } - } - - for (uint32_t iter = 0; iter < num_iters; iter++) { - const auto buf_idx = iter % kNumStages; - const auto offset = iter * kSizePerStage; - const auto this_size = min(kSizePerStage, length - offset); - - if (lane_id == 1) { - const auto phase_bit = (iter / kNumStages) & 1; - ptx::mbarrier_wait(&smem->barrier[kIsScatter][buf_idx], phase_bit); - } - __syncwarp(); - -#pragma unroll - for (uint32_t i = 0; i < kElemPerStage; i++) { - const auto local_idx = tx + i * kBlockSize; - if (local_idx >= this_size) break; - const auto score = smem->score_buffer[buf_idx][local_idx]; - const auto bin = extract_coarse_bin(score); - if constexpr (kIsScatter) { - const auto global_idx = offset + local_idx; - if (bin > thr_bin) { - const auto pos = atomicAdd(&smem->counter_gt, 1); - if (pos < K) s_topk_indices[pos] = global_idx; - } else if (bin == thr_bin) { - const auto pos = atomicAdd(&smem->counter_eq, 1); - if (pos < kMaxTies) smem->tie_buffer[pos] = {global_idx, score}; - } - } else { - atomicAdd(&smem->histogram[bin], 1); - } - } - - __syncthreads(); - if (tx == 0) { - if (const auto next_iter = iter + kNumStages; next_iter < num_iters) { - issue_tma(scores, next_iter, length_aligned, smem); - } - } - } - } - - // --------------------------------------------------------------------------- - // Phase B: find the threshold bin via a warp-level prefix scan. - // Same structure as SmallTopKImpl's phase 2 (4 bins/thread, warp_sum relay). - // --------------------------------------------------------------------------- - - SGL_DEVICE static void find_threshold(uint32_t length, Smem* smem) { - const auto tx = threadIdx.x; - const auto lane_id = tx % kWarpThreads; - const auto warp_id = tx / kWarpThreads; - - uint32_t orig[kHistItems]; - const auto hist_vec = smem->histogram_vec[tx]; - uint32_t local_sum = 0; -#pragma unroll - for (uint32_t i = 0; i < kHistItems; ++i) { - orig[i] = hist_vec[i]; - local_sum += orig[i]; - } - - const auto warp_inc = warp_inclusive_sum(lane_id, local_sum); - const auto warp_exc = warp_inc - local_sum; - if (lane_id == kWarpThreads - 1) smem->warp_sum[warp_id] = warp_inc; - __syncthreads(); - - const auto tmp = smem->warp_sum[lane_id]; - uint32_t prefix_sum = warp::reduce_sum(lane_id < warp_id ? tmp : 0); - prefix_sum += warp_exc; -#pragma unroll - for (uint32_t i = 0; i < kHistItems; ++i) { - prefix_sum += orig[i]; - const auto above = length - prefix_sum; - if (above < K && above + orig[i] >= K) { - smem->match = { - .bin = tx * kHistItems + i, - .above_count = above, - .equal_count = orig[i], - }; - } - } - __syncthreads(); - } - - SGL_DEVICE static void run(const float* scores, const uint32_t length, int32_t* topk_indices, void* _smem) { - const auto smem = static_cast(_smem); - const auto tx = threadIdx.x; - __builtin_assume(tx < kBlockSize); - - // Init histogram, barriers, counters. - { - HistVec zero; - zero.fill(0); - smem->histogram_vec[tx] = zero; - if (tx < 2 * kNumStages) { - const auto base_barrier = &smem->barrier[0][0]; - ptx::mbarrier_init(&base_barrier[tx], 1); - } - if (tx == 0) { - smem->counter_gt = 0; - smem->counter_eq = 0; - } - __syncthreads(); - } - - // Phase A: histogram pass (pipelined TMA stream). - stream_pass(scores, length, 0, nullptr, smem); - - // Phase B: locate threshold bin & re-init barriers - find_threshold(length, smem); - - // Phase C: scatter pass. - stream_pass(scores, length, smem->match.bin, topk_indices, smem); - } - - SGL_DEVICE static void transform(const TransformParams params, void* _smem) { - // Phase D: page-translate above entries, then refine ties. - const auto smem = static_cast(_smem); - const auto tx = threadIdx.x; - const auto num_above = smem->match.above_count; - if (tx < num_above) params.transform(tx); - const auto num_equal = smem->counter_eq; - if (num_above >= K || num_equal == 0) return; - const auto clamped_ties = min(num_equal, kMaxTies); - tie_handle_transform(smem->tie_buffer, clamped_ties, num_above, K, params, &smem->stage2); - } -}; - -} // namespace device::top512 diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh deleted file mode 100644 index e0ce2dc086..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/common.cuh +++ /dev/null @@ -1,120 +0,0 @@ -#pragma once -#include - -namespace device::distributed { - -inline constexpr uint32_t kMaxNumGPU = 8; - -struct alignas(128) Semaphore { - public: - constexpr Semaphore() : m_flag(0), m_counter(0) {} - - template - SGL_DEVICE uint32_t get() const { - uint32_t val; - if constexpr (kFence) { - asm volatile("ld.acquire.sys.global.u32 %0, [%1];" : "=r"(val) : "l"(&m_flag)); - } else { - asm volatile("ld.volatile.global.u32 %0, [%1];" : "=r"(val) : "l"(&m_flag)); - } - return val; - } - - template - SGL_DEVICE uint32_t add(uint32_t val) { - uint32_t old_val; - if constexpr (kFence) { - asm volatile("atom.release.sys.global.add.u32 %0, [%1], %2;" : "=r"(old_val) : "l"(&m_flag), "r"(val)); - } else { - asm volatile("atom.global.add.u32 %0, [%1], %2;" : "=r"(old_val) : "l"(&m_flag), "r"(val)); - } - return old_val; - } - - // Only called by the owning GPU - plain load is sufficient - SGL_DEVICE uint32_t get_counter() const { - return m_counter; - } - - // Only called by the owning GPU - plain store is sufficient - SGL_DEVICE void set_counter(uint32_t val) { - m_counter = val; - } - - private: - uint32_t m_flag; - uint32_t m_counter; -}; - -struct PullController { - public: - using SignalType = Semaphore; - - PullController(void** signals, uint32_t num_gpu) { - for (uint32_t i = 0; i < num_gpu; ++i) { - m_signals[i] = static_cast(signals[i]); - } - } - - /// Synchronize all GPUs. - /// When kFence is true, establishes happens-before across GPUs using - /// release/acquire semantics, ensuring prior writes are visible system-wide. - template - SGL_DEVICE void sync(uint32_t rank, uint32_t num_gpu) const { - // For fenced sync: ensure all threads in this block have completed their writes, - // so the signaling thread's release carries them transitively. - static_assert(!(kFence && kStart), "Start stage does not need to wait fence"); - if constexpr (kFence || !kStart) __syncthreads(); - constexpr auto kStage = kStart ? 1 : 2; - const auto warp_id = threadIdx.x / kWarpThreads; - const auto lane_id = threadIdx.x % kWarpThreads; - if (lane_id == 0 && warp_id < num_gpu) { - auto& signal = m_signals[warp_id][blockIdx.x]; - signal.add(1); - if (warp_id == rank) { - const auto target = num_gpu * kStage; - /// NOTE: correctness here: - /// - base is only read/updated locally by the owning GPU - const auto base = signal.get_counter(); - while (signal.get() - base < target) - ; - if constexpr (!kStart) { - signal.set_counter(base + target); - } - } - } - if constexpr (kStart) __syncthreads(); - } - - private: - Semaphore* __restrict__ m_signals[kMaxNumGPU]; -}; - -struct PushController { - public: - using SignalType = uint32_t; - static constexpr int64_t kNumStages = 2; - - PushController(void* ptr) : m_local_signal(static_cast(ptr)) {} - - SGL_DEVICE SignalType epoch() const { - return m_local_signal[blockIdx.x]; - } - - SGL_DEVICE void exit() const { - __syncthreads(); - if (threadIdx.x == 0) { - this->exit_unsafe(blockIdx.x); - } - } - - SGL_DEVICE void exit_unsafe(uint32_t which) const { - auto& signal = m_local_signal[which]; - signal = (signal + 1) % kNumStages; - } - - private: - SignalType* m_local_signal; -}; - -} // namespace device::distributed diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh deleted file mode 100644 index 239fac71a1..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/distributed/custom_all_reduce.cuh +++ /dev/null @@ -1,354 +0,0 @@ -#pragma once -#include - -#include -#include - -#include - -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -namespace host::distributed { - -using device::distributed::PullController, device::distributed::PushController; - -struct AllReduceData { - constexpr AllReduceData() {} - void* __restrict__ input[device::distributed::kMaxNumGPU]; -}; - -using ExternHandle = tvm::ffi::Array; - -inline ExternHandle to_extern_handle(void* ptr) { - ExternHandle array; - cudaIpcMemHandle_t handle; - RuntimeDeviceCheck(cudaIpcGetMemHandle(&handle, ptr)); - for (size_t i = 0; i < sizeof(handle); ++i) { - array.push_back(handle.reserved[i]); - } - return array; -} - -inline void* from_extern_handle(const ExternHandle& array) { - cudaIpcMemHandle_t handle; - RuntimeCheck(array.size() == sizeof(handle), "Invalid IPC handle size: ", array.size()); - for (size_t i = 0; i < sizeof(handle); ++i) { - handle.reserved[i] = array[i]; - } - void* ptr; - RuntimeDeviceCheck(cudaIpcOpenMemHandle(&ptr, handle, cudaIpcMemLazyEnablePeerAccess)); - return ptr; -} - -struct HandleHash { - std::size_t operator()(const cudaIpcMemHandle_t& handle) const { - return std::hash{}({handle.reserved, sizeof(handle.reserved)}); - } -}; - -struct HandleEqual { - bool operator()(const cudaIpcMemHandle_t& a, const cudaIpcMemHandle_t& b) const { - return std::memcmp(a.reserved, b.reserved, sizeof(a.reserved)) == 0; - } -}; - -/** - * \brief The control plane of the custom all-reduce implementation. - * It manages the internal state and synchronization of the participating GPUs. - */ -struct CustomAllReduceBase : public tvm::ffi::Object { - public: - TVM_FFI_DECLARE_OBJECT_INFO_FINAL("sgl.CustomAllReduce", CustomAllReduceBase, tvm::ffi::Object); - - static constexpr bool _type_mutable = true; - using InputPair = tvm::ffi::Tuple; // (offset, ipc handle) - - CustomAllReduceBase( - uint32_t rank, - uint32_t num_gpu, - uint32_t max_num_cta_pull, - uint32_t max_num_cta_push, - int64_t pull_buffer_size, - int64_t push_buffer_size, - int64_t graph_buffer_count) - : m_pull_buffer_bytes(pull_buffer_size), - m_push_buffer_bytes(push_buffer_size), - m_graph_buffer_count(graph_buffer_count), - m_rank(rank), - m_num_gpu(num_gpu), - m_max_num_cta_pull(max_num_cta_pull), - m_max_num_cta_push(max_num_cta_push), - // default config for pull kernel, can be updated by `configure()` - m_num_cta(max_num_cta_pull), - m_cta_size(256) { - RuntimeCheck(pull_buffer_size % 128 == 0, "Pull buffer size should be aligned to 128 bytes"); - RuntimeCheck(push_buffer_size % 128 == 0, "Push buffer size should be aligned to 128 bytes"); - RuntimeCheck(rank < num_gpu, "Invalid rank: ", rank); - const int64_t kU32Max = static_cast(std::numeric_limits::max()); - const int64_t push_buffer_size_all = push_all_ranks_bytes(); - RuntimeCheck(pull_buffer_size <= kU32Max, "Pull buffer size is too large: ", pull_buffer_size); - RuntimeCheck(push_buffer_size_all <= kU32Max, "Push buffer size is too large: ", push_buffer_size_all); - RuntimeDeviceCheck(cudaMalloc(&m_storage, storage_bytes())); - } - - ExternHandle share_storage() { - return to_extern_handle(m_storage); - } - - tvm::ffi::Array share_graph_inputs() { - tvm::ffi::Array result; - const auto new_inputs_count = registered_count() - m_cum_registered_count; - RuntimeCheck(new_inputs_count >= 0, "Invalid new count: ", new_inputs_count); - result.reserve(new_inputs_count); - std::unordered_map ipc_cache; - const auto get_handle = [&](void* ptr) -> ExternHandle { - const auto it = ipc_cache.find(ptr); - if (it != ipc_cache.end()) return it->second; - const auto handle = to_extern_handle(ptr); - ipc_cache.try_emplace(ptr, handle); - return handle; - }; - for (const auto ptr : std::span(m_graph_capture_inputs).subspan(m_cum_registered_count)) { - // note: must share the base address of each allocation, or we get wrong address - void* base_ptr; - const auto cu_result = cuPointerGetAttribute(&base_ptr, CU_POINTER_ATTRIBUTE_RANGE_START_ADDR, (CUdeviceptr)ptr); - RuntimeCheck(cu_result == CUDA_SUCCESS, "failed to get pointer attr"); - const auto offset = reinterpret_cast(ptr) - reinterpret_cast(base_ptr); - result.push_back(InputPair{offset, get_handle(base_ptr)}); - } - return result; - } - - void post_init(tvm::ffi::Array ipc_storages) { - RuntimeCheck(ipc_storages.size() == m_num_gpu, "Invalid array size: ", ipc_storages.size()); - m_peer_storage.resize(m_num_gpu); - for (const auto i : irange(m_num_gpu)) { - if (i == m_rank) { - m_peer_storage[i] = m_storage; - } else { - m_peer_storage[i] = from_extern_handle(ipc_storages[i]); - } - } - - // set signal buffer to zero - const auto pull_signal = get_pull_signal(m_storage); - RuntimeDeviceCheck(cudaMemset(pull_signal, 0, pull_signal_bytes())); - - // update the pull controller and data pointer - RuntimeCheck(!m_pull_ctrl.has_value(), "Controller is already initialized"); - m_pull_ctrl.emplace(m_peer_storage.data(), m_num_gpu); - AllReduceData data; - for (const auto i : irange(m_num_gpu)) { - data.input[i] = get_pull_buffer(m_peer_storage[i]); - } - const auto default_data_ptr = get_data_ptr(); - RuntimeDeviceCheck(cudaMemcpy(default_data_ptr, &data, sizeof(AllReduceData), cudaMemcpyHostToDevice)); - - // update the push controller and data pointer - RuntimeCheck(!m_push_ctrl.has_value(), "Controller is already initialized"); - const auto push_signal = get_push_signal(m_storage); - RuntimeDeviceCheck(cudaMemset(push_signal, 0, push_signal_bytes())); - m_push_ctrl.emplace(push_signal); - const auto push_buffer = get_push_buffer(m_storage); - RuntimeDeviceCheck(cudaMemset(push_buffer, 0, push_all_ranks_bytes())); - } - - void register_inputs(tvm::ffi::Array> ipc_graph_inputs) { - RuntimeCheck(ipc_graph_inputs.size() == m_num_gpu); - const auto new_registered_count = registered_count() - m_cum_registered_count; - RuntimeCheck(new_registered_count >= 0, "Invalid registered count: ", new_registered_count); - if (new_registered_count == 0) return; // avoid `m_get_data_ptr()` out-of-bounds - std::vector data; - data.resize(new_registered_count); - const auto open_cached = [&](const ExternHandle& h) -> void* { - RuntimeCheck(h.size() == sizeof(cudaIpcMemHandle_t), "Invalid IPC handle size: ", h.size()); - cudaIpcMemHandle_t handle; - for (size_t i = 0; i < sizeof(handle); ++i) - handle.reserved[i] = h[i]; - const auto [it, success] = m_ipc_cache.try_emplace(handle, nullptr); - if (success) { - void* ptr; - RuntimeDeviceCheck(cudaIpcOpenMemHandle(&ptr, handle, cudaIpcMemLazyEnablePeerAccess)); - it->second = ptr; - } - return it->second; - }; - for (const auto i : irange(ipc_graph_inputs.size())) { - const auto& array = ipc_graph_inputs[i]; - RuntimeCheck(int64_t(array.size()) == new_registered_count); - if (i == m_rank) { - for (const auto j : irange(new_registered_count)) { - data[j].input[i] = m_graph_capture_inputs[m_cum_registered_count + j]; - } - } else { - for (const auto j : irange(new_registered_count)) { - /// NOTE: structural binding will cause intern compiler error... - const auto elem = array[j]; - const auto offset = elem.get<0>(); - const auto ipc_handle = elem.get<1>(); - data[j].input[i] = pointer::offset(open_cached(ipc_handle), offset); - } - } - } - - const auto new_registered_bytes = sizeof(AllReduceData) * new_registered_count; - const auto dst_ptr = get_data_ptr(m_cum_registered_count); - m_cum_registered_count += new_registered_count; - RuntimeDeviceCheck(cudaMemcpy(dst_ptr, data.data(), new_registered_bytes, cudaMemcpyHostToDevice)); - } - - void set_cuda_graph_capture(bool enabled) { - m_is_graph_capturing = enabled; - } - - void free_ipc_handles() { - for (const auto& pair : m_ipc_cache) { - host::RuntimeDeviceCheck(cudaIpcCloseMemHandle(pair.second)); - } - m_ipc_cache.clear(); - } - - void free_storage() { - host::RuntimeDeviceCheck(cudaFree(m_storage)); - m_storage = nullptr; - } - - tvm::ffi::Tuple configure_pull(uint32_t num_cta, uint32_t cta_size) { - using host::RuntimeCheck; - const auto min_cta_size = m_num_gpu * device::kWarpThreads; - RuntimeCheck(num_cta > 0 && num_cta <= m_max_num_cta_pull, "Invalid number of CTAs: ", num_cta); - RuntimeCheck(cta_size >= min_cta_size, "Block size must be at least ", min_cta_size); - const auto old_num_cta = m_num_cta; - const auto old_block_size = m_cta_size; - m_num_cta = num_cta; - m_cta_size = cta_size; - return tvm::ffi::Tuple{old_num_cta, old_block_size}; - } - - protected: - AllReduceData* allocate_graph_capture_input(void* data_ptr) { - const auto count = registered_count(); - RuntimeCheck(count < m_graph_buffer_count, "Graph buffer overflow, increase `graph_buffer_count`!"); - m_graph_capture_inputs.push_back(data_ptr); - return get_data_ptr(count); - } - AllReduceData* get_data_ptr(int64_t which = -1) { - const auto count = registered_count(); - RuntimeCheck(which >= -1 && which < count, "Invalid graph buffer index: ", which, ", count: ", count); - const auto start = get_pull_params(m_storage); - return static_cast(start) + (1 + which); - } - int64_t registered_count() const { - return static_cast(m_graph_capture_inputs.size()); - } - int64_t pull_signal_bytes() const { - return _align_bytes(sizeof(PullController::SignalType) * m_max_num_cta_pull); - } - int64_t push_signal_bytes() const { - return _align_bytes(sizeof(PushController::SignalType) * m_max_num_cta_push); - } - int64_t graph_param_bytes() const { - return _align_bytes(sizeof(AllReduceData) * (1 + m_graph_buffer_count)); // 1 for default - } - int64_t push_all_ranks_bytes() const { - return _align_bytes(PushController::kNumStages * m_num_gpu * m_push_buffer_bytes); - } - int64_t storage_bytes() const { - return _get_offset_impl(5); - } - void* get_pull_signal(void* ptr) const { - return pointer::offset(ptr, _get_offset_impl(0)); - } - void* get_push_signal(void* ptr) const { - return pointer::offset(ptr, _get_offset_impl(1)); - } - void* get_pull_params(void* ptr) const { - return pointer::offset(ptr, _get_offset_impl(2)); - } - void* get_pull_buffer(void* ptr) const { - return pointer::offset(ptr, _get_offset_impl(3)); - } - void* get_push_buffer(void* ptr) const { - return pointer::offset(ptr, _get_offset_impl(4)); - } - int64_t _get_offset_impl(int64_t which) const { - // | SignalArray (pull + push) | GraphBuffers (pull params) | Buffers (pull + push) | - const int64_t offset_map[5] = { - /*[0]=*/pull_signal_bytes(), - /*[1]=*/push_signal_bytes(), - /*[2]=*/graph_param_bytes(), - /*[3]=*/m_pull_buffer_bytes, - /*[4]=*/push_all_ranks_bytes(), - }; - RuntimeCheck(which >= 0 && which <= 5, "Invalid offset index: ", which); - return std::accumulate(offset_map, offset_map + which, int64_t(0)); - } - static int64_t _align_bytes(int64_t size) { - return div_ceil(size, 128) * 128; - } - - const int64_t m_pull_buffer_bytes; - const int64_t m_push_buffer_bytes; - const int64_t m_graph_buffer_count; - const uint32_t m_rank; - const uint32_t m_num_gpu; - const uint32_t m_max_num_cta_pull; - const uint32_t m_max_num_cta_push; - // these 2 config should only affect pull kernel - uint32_t m_num_cta; - uint32_t m_cta_size; - // other states - bool m_is_graph_capturing = false; - int64_t m_cum_registered_count = 0; - std::optional m_pull_ctrl; - std::optional m_push_ctrl; - void* m_storage = nullptr; - std::vector m_graph_capture_inputs; - std::vector m_peer_storage; - std::unordered_map m_ipc_cache; -}; - -struct CustomAllReduceRef : public tvm::ffi::ObjectRef { - TVM_FFI_DEFINE_OBJECT_REF_METHODS_NOTNULLABLE(CustomAllReduceRef, tvm::ffi::ObjectRef, CustomAllReduceBase); -}; - -} // namespace host::distributed - -namespace device::distributed { - -template -SGL_DEVICE auto reduce_impl(AlignedVector (&storage)[M]) -> AlignedVector { - fp32x2_t acc[N] = {}; -#pragma unroll // unroll num gpu - for (uint32_t i = 0; i < M; ++i) { -#pragma unroll // unroll vec - for (uint32_t j = 0; j < N; ++j) { - const auto [x, y] = cast(storage[i][j]); - auto& [x_acc, y_acc] = acc[j]; - x_acc += x; - y_acc += y; - } - } - - AlignedVector result; -#pragma unroll - for (uint32_t j = 0; j < N; ++j) { - result[j] = cast(acc[j]); - } - - return result; -} - -} // namespace device::distributed diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h deleted file mode 100644 index 17d9048d4c..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/ffi.h +++ /dev/null @@ -1,104 +0,0 @@ -#pragma once -#include - -#include -#include -#include -#include - -#include -#include -#include -#include -#include - -namespace host::ffi { - -using tvm::ffi::Tensor, tvm::ffi::TensorView, tvm::ffi::ShapeView; - -inline Tensor empty(ShapeView shape, DLDataType dtype, DLDevice device) { - return Tensor::FromEnvAlloc(::TVMFFIEnvTensorAlloc, shape, dtype, device); -} - -inline Tensor empty_like(TensorView tensor) { - return empty(tensor.shape(), tensor.dtype(), tensor.device()); -} - -struct _dummy_deleter { - void operator()(void*) const {} -}; - -// template - -template -struct FromBlobContext { - [[no_unique_address]] Fn deleter; - int64_t dimension; - int64_t* get_shape() { - return reinterpret_cast(this + 1); - } - int64_t* get_stride() { - return this->get_shape() + dimension; - } -}; - -template -inline Tensor from_blob( - void* data, - ShapeView shape, - DLDataType dtype, - DLDevice device, - Fn&& deleter = {}, - std::optional stride = {}, - uint64_t byte_offset = 0) { - using Context = FromBlobContext>; - const auto ndim = shape.size(); - const auto ctx = [&] { - auto ptr = std::malloc(sizeof(Context) + sizeof(int64_t) * ndim * 2); - auto ctx = static_cast(ptr); - std::construct_at(ctx, std::forward(deleter), static_cast(ndim)); - stdr::copy_n(shape.data(), ndim, ctx->get_shape()); - if (stride.has_value()) { - RuntimeCheck(stride->size() == ndim, "Stride ndim mismatch!"); - stdr::copy_n(stride->data(), ndim, ctx->get_stride()); - } else { - int64_t stride_val = 1; - for (const auto i : irange(ndim)) { - const auto j = ndim - 1 - i; - ctx->get_stride()[j] = stride_val; - stride_val *= shape[j]; - } - } - return ctx; - }(); - const auto tensor = DLTensor{ - .data = data, - .device = device, - .ndim = static_cast(ndim), - .dtype = dtype, - .shape = ctx->get_shape(), - .strides = ctx->get_stride(), - .byte_offset = byte_offset, - }; - const auto blob_deleter = [](DLManagedTensor* self) { - auto ctx = static_cast(self->manager_ctx); - ctx->deleter(self->dl_tensor.data); - std::destroy_at(ctx); - std::free(ctx); - }; - auto managed_tensor = DLManagedTensor{tensor, ctx, blob_deleter}; - return Tensor::FromDLPack(&managed_tensor); -} - -template -inline Tensor from_blob_like( - void* data, - TensorView t, - Fn&& deleter = {}, - bool is_contiguous = false, // if override to true, the stride will be ignored - uint64_t byte_offset = 0) { - const auto stride = is_contiguous ? std::nullopt : std::optional{t.strides()}; - return from_blob(data, t.shape(), t.dtype(), t.device(), std::forward(deleter), stride, byte_offset); -} - -} // namespace host::ffi diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh deleted file mode 100644 index cd024acd46..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/impl/norm.cuh +++ /dev/null @@ -1,168 +0,0 @@ -#pragma once -#include -#include -#include -#include -#include - -#include -#include - -namespace host::norm { - -/** - * \brief Check if the given configuration is supported. - * \tparam T Element type (only fp16_t/bf16_t is supported) - * \tparam kDim Dimension size (usually hidden size) - */ -template -inline constexpr bool is_config_supported() { - if (!std::is_same_v && !std::is_same_v) return false; - if (kDim <= 256) { - return (kDim == 64 || kDim == 128 || kDim == 256); - } else { - return (kDim % 256 == 0 && kDim <= 8192); - } -} - -/** - * \brief Determine whether to use cta norm based on dimension size. - * TL;DR: use warp norm for dim <= 256, cta norm otherwise. - * \tparam T Element type (fp16_t or bf16_t) - * \tparam kDim Dimension size (usually hidden size) - * \note This function assumes that the configuration is supported. - * \see `is_config_supported` - */ -template -inline constexpr bool should_use_cta() { - static_assert(is_config_supported(), "Unsupported norm configuration"); - return kDim > 256; -} - -/** - * \brief Get the number of threads per CTA for cta norm. - * \tparam T Element type (fp16_t or bf16_t) - * \tparam kDim Dimension size (usually hidden size) - * \return Number of threads per CTA - */ -template -inline constexpr uint32_t get_cta_threads() { - static_assert(should_use_cta()); - return (kDim / 256) * device::kWarpThreads; -} - -} // namespace host::norm - -namespace device::norm { - -namespace details { - -template -SGL_DEVICE AlignedVector apply_norm_impl( - const AlignedVector input, - const AlignedVector weight, - const float eps, - [[maybe_unused]] float* smem_buffer, - [[maybe_unused]] uint32_t num_warps) { - float sum_of_squares = 0.0f; - -#pragma unroll - for (auto i = 0u; i < N; ++i) { - const auto fp32_input = cast(input[i]); - sum_of_squares += fp32_input.x * fp32_input.x; - sum_of_squares += fp32_input.y * fp32_input.y; - } - - sum_of_squares = warp::reduce_sum(sum_of_squares); - float norm_factor; - if constexpr (kUseCTA) { - // need to synchronize across the cta - const auto warp_id = threadIdx.x / kWarpThreads; - smem_buffer[warp_id] = sum_of_squares; - __syncthreads(); - // use the first warp to reduce - if (warp_id == 0) { - const auto tx = threadIdx.x; - const auto local_sum = tx < num_warps ? smem_buffer[tx] : 0.0f; - sum_of_squares = warp::reduce_sum(local_sum); - smem_buffer[32] = math::rsqrt(sum_of_squares / kDim + eps); - } - __syncthreads(); - norm_factor = smem_buffer[32]; - } else { - norm_factor = math::rsqrt(sum_of_squares / kDim + eps); - } - - AlignedVector output; - -#pragma unroll - for (auto i = 0u; i < N; ++i) { - const auto fp32_input = cast(input[i]); - const auto fp32_weight = cast(weight[i]); - output[i] = cast({ - fp32_input.x * norm_factor * fp32_weight.x, - fp32_input.y * norm_factor * fp32_weight.y, - }); - } - - return output; -} - -} // namespace details - -/** - * \brief Apply norm using warp-level implementation. - * \tparam kDim Dimension size - * \tparam T Element type (fp16_t or bf16_t) - * \param input Input vector - * \param weight Weight vector - * \param eps Epsilon value for numerical stability - * \return Normalized output vector - */ -template -SGL_DEVICE T apply_norm_warp(const T& input, const T& weight, float eps) { - static_assert(kDim <= 256, "Warp norm only supports dim <= 256"); - return details::apply_norm_impl(input, weight, eps, nullptr, 0); -} - -/** - * \brief Apply norm using CTA-level implementation. - * \tparam kDim Dimension size - * \tparam T Element type (fp16_t or bf16_t) - * \param input Input vector - * \param weight Weight vector - * \param eps Epsilon value for numerical stability - * \param smem Shared memory buffer - * \param num_warps Number of warps in the CTA - * \return Normalized output vector - */ -template -SGL_DEVICE T apply_norm_cta( - const T& input, const T& weight, float eps, float* smem, uint32_t num_warps = blockDim.x / kWarpThreads) { - static_assert(kDim > 256, "CTA norm only supports dim > 256"); - return details::apply_norm_impl(input, weight, eps, smem, num_warps); -} - -/** - * \brief Storage type for norm operation. - * For warp norm, the storage size depends on kDim. - * For cta norm, the storage size is fixed to 16B. - * We will also pack the input 16-bit floats into 32-bit types - * for faster CUDA core operations. - * - * \tparam T Element type (fp16_t or bf16_t) - * \tparam kDim Dimension size - */ -template -using StorageType = std::conditional_t< // storage type - (kDim > 256), // whether to use cta norm - AlignedVector, 4>, // cta norm storage, fixed to 16B - AlignedVector, kDim / (2 * kWarpThreads)> // warp norm storage - >; - -/** - * \brief Minimum shared memory size (in bytes) required for cta norm. - */ -inline constexpr uint32_t kSmemBufferSize = 33; - -} // namespace device::norm diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh deleted file mode 100644 index 4f9ac48141..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/math.cuh +++ /dev/null @@ -1,71 +0,0 @@ -/// \file math.cuh -/// \brief Device-side math helper functions and constants. -/// -/// Provides type-generic wrappers around CUDA math intrinsics by -/// dispatching through `dtype_trait`. All functions are forced-inline -/// device functions. - -#pragma once -#include - -#include - -namespace device::math { - -/// \brief Constant: log2(e) -inline constexpr float log2e = 1.44269504088896340736f; -/// \brief Constant: ln(2) -inline constexpr float loge2 = 0.693147180559945309417f; -/// \brief Maximum representable value for FP8 E4M3 format. -inline constexpr float FP8_E4M3_MAX = 448.0f; -static_assert(log2e * loge2 == 1.0f, "log2e * loge2 must be 1"); - -/// \brief Returns the larger of `a` and `b`. -template -SGL_DEVICE T max(T a, T b) { - return dtype_trait::max(a, b); -} - -/// \brief Returns the smaller of `a` and `b`. -template -SGL_DEVICE T min(T a, T b) { - return dtype_trait::min(a, b); -} - -/// \brief Returns the absolute value of `a`. -template -SGL_DEVICE T abs(T a) { - return dtype_trait::abs(a); -} - -/// \brief Returns the square root of `a`. -template -SGL_DEVICE T sqrt(T a) { - return dtype_trait::sqrt(a); -} - -/// \brief Returns the reciprocal square root of `a` (i.e. 1 / sqrt(a)). -template -SGL_DEVICE T rsqrt(T a) { - return dtype_trait::rsqrt(a); -} - -/// \brief Returns e^a. -template -SGL_DEVICE T exp(T a) { - return dtype_trait::exp(a); -} - -/// \brief Returns sin(a). -template -SGL_DEVICE T sin(T a) { - return dtype_trait::sin(a); -} - -/// \brief Returns cos(a). -template -SGL_DEVICE T cos(T a) { - return dtype_trait::cos(a); -} - -} // namespace device::math diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh deleted file mode 100644 index 4ea722a3fe..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/runtime.cuh +++ /dev/null @@ -1,86 +0,0 @@ -/// \file runtime.cuh -/// \brief Host-side CUDA runtime query helpers. -/// -/// Thin wrappers around CUDA occupancy and device-property APIs with -/// automatic error checking via `RuntimeDeviceCheck`. - -#pragma once - -#include - -#include -#include -#ifndef USE_ROCM -#include -#else -#include -#ifndef cudaOccupancyMaxActiveBlocksPerMultiprocessor -#define cudaOccupancyMaxActiveBlocksPerMultiprocessor hipOccupancyMaxActiveBlocksPerMultiprocessor -#endif -#ifndef cudaDeviceGetAttribute -#define cudaDeviceGetAttribute hipDeviceGetAttribute -#endif -#ifndef cudaDevAttrMultiProcessorCount -#define cudaDevAttrMultiProcessorCount hipDeviceAttributeMultiprocessorCount -#endif -#ifndef cudaDevAttrComputeCapabilityMajor -#define cudaDevAttrComputeCapabilityMajor hipDeviceAttributeComputeCapabilityMajor -#endif -#ifndef cudaRuntimeGetVersion -#define cudaRuntimeGetVersion hipRuntimeGetVersion -#endif -#ifndef cudaOccupancyAvailableDynamicSMemPerBlock -inline hipError_t -cudaOccupancyAvailableDynamicSMemPerBlock(std::size_t* smem, const void* func, int num_blocks, int block_size) { - // HIP does not expose this directly; return max shared mem as conservative estimate - hipDeviceProp_t prop; - int device; - hipGetDevice(&device); - hipGetDeviceProperties(&prop, device); - *smem = prop.sharedMemPerBlock; - return hipSuccess; -} -#endif -#endif - -namespace host::runtime { - -// Return the maximum number of active blocks per SM for the given kernel -template -inline auto get_blocks_per_sm(T&& kernel, int32_t block_dim, std::size_t dynamic_smem = 0) -> uint32_t { - int num_blocks_per_sm = 0; - RuntimeDeviceCheck( - cudaOccupancyMaxActiveBlocksPerMultiprocessor(&num_blocks_per_sm, kernel, block_dim, dynamic_smem)); - return static_cast(num_blocks_per_sm); -} - -// Return the number of SMs for the given device -inline auto get_sm_count(int device_id) -> uint32_t { - int sm_count; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&sm_count, cudaDevAttrMultiProcessorCount, device_id)); - return static_cast(sm_count); -} - -// Return the Major compute capability for the given device -inline auto get_cc_major(int device_id) -> int { - int cc_major; - RuntimeDeviceCheck(cudaDeviceGetAttribute(&cc_major, cudaDevAttrComputeCapabilityMajor, device_id)); - return cc_major; -} - -// Return the runtime version -inline auto get_runtime_version() -> int { - int runtime_version; - RuntimeDeviceCheck(cudaRuntimeGetVersion(&runtime_version)); - return runtime_version; -} - -// Return the maximum dynamic shared memory per block for the given kernel -template -inline auto get_available_dynamic_smem_per_block(T&& kernel, int num_blocks, int block_size) -> std::size_t { - std::size_t smem_size; - RuntimeDeviceCheck(cudaOccupancyAvailableDynamicSMemPerBlock(&smem_size, kernel, num_blocks, block_size)); - return smem_size; -} - -} // namespace host::runtime diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp b/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp deleted file mode 100644 index d229d3a975..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/scalar_type.hpp +++ /dev/null @@ -1,334 +0,0 @@ -#pragma once - -#include -#include -#ifndef __CUDACC__ -#include -#endif - -namespace host { - -// -// ScalarType can represent a wide range of floating point and integer types, -// in particular it can be used to represent sub-byte data types (something -// that torch.dtype currently does not support). -// -// The type definitions on the Python side can be found in: vllm/scalar_type.py -// these type definitions should be kept up to date with any Python API changes -// here. -// -class ScalarType { - public: - enum NanRepr : uint8_t { - NAN_NONE = 0, // nans are not supported - NAN_IEEE_754 = 1, // nans are: exp all 1s, mantissa not all 0s - NAN_EXTD_RANGE_MAX_MIN = 2, // nans are: exp all 1s, mantissa all 1s - - NAN_REPR_ID_MAX - }; - - constexpr ScalarType( - uint8_t exponent, - uint8_t mantissa, - bool signed_, - int32_t bias, - bool finite_values_only = false, - NanRepr nan_repr = NAN_IEEE_754) - : exponent(exponent), - mantissa(mantissa), - signed_(signed_), - bias(bias), - finite_values_only(finite_values_only), - nan_repr(nan_repr) {}; - - static constexpr ScalarType int_(uint8_t size_bits, int32_t bias = 0) { - return ScalarType(0, size_bits - 1, true, bias); - } - - static constexpr ScalarType uint(uint8_t size_bits, int32_t bias = 0) { - return ScalarType(0, size_bits, false, bias); - } - - // IEEE 754 compliant floating point type - static constexpr ScalarType float_IEEE754(uint8_t exponent, uint8_t mantissa) { - assert(mantissa > 0 && exponent > 0); - return ScalarType(exponent, mantissa, true, 0, false, NAN_IEEE_754); - } - - // IEEE 754 non-compliant floating point type - static constexpr ScalarType float_(uint8_t exponent, uint8_t mantissa, bool finite_values_only, NanRepr nan_repr) { - assert(nan_repr < NAN_REPR_ID_MAX); - assert(mantissa > 0 && exponent > 0); - assert(nan_repr != NAN_IEEE_754); - return ScalarType(exponent, mantissa, true, 0, finite_values_only, nan_repr); - } - - uint8_t const exponent; // size of the exponent field (0 for integer types) - uint8_t const mantissa; // size of the mantissa field (size of the integer - // excluding the sign bit for integer types) - bool const signed_; // flag if the type supports negative numbers (i.e. has a - // sign bit) - int32_t const bias; // stored values equal value + bias, - // used for quantized type - - // Extra Floating point info - bool const finite_values_only; // i.e. no +/-inf if true - NanRepr const nan_repr; // how NaNs are represented - // (not applicable for integer types) - - using Id = int64_t; - - private: - // Field size in id - template - static constexpr size_t member_id_field_width() { - using T = std::decay_t; - return std::is_same_v ? 1 : sizeof(T) * 8; - } - - template - static constexpr auto reduce_members_helper(Fn f, Init val, Member member, Rest... rest) { - auto new_val = f(val, member); - if constexpr (sizeof...(rest) > 0) { - return reduce_members_helper(f, new_val, rest...); - } else { - return new_val; - }; - } - - template - constexpr auto reduce_members(Fn f, Init init) const { - // Should be in constructor order for `from_id` - return reduce_members_helper(f, init, exponent, mantissa, signed_, bias, finite_values_only, nan_repr); - }; - - template - static constexpr auto reduce_member_types(Fn f, Init init) { - constexpr auto dummy_type = ScalarType(0, 0, false, 0, false, NAN_NONE); - return dummy_type.reduce_members(f, init); - }; - - static constexpr auto id_size_bits() { - return reduce_member_types( - [](int acc, auto member) -> int { return acc + member_id_field_width(); }, 0); - } - - public: - // unique id for this scalar type that can be computed at compile time for - // c++17 template specialization this is not needed once we migrate to - // c++20 and can pass literal classes as template parameters - constexpr Id id() const { - static_assert(id_size_bits() <= sizeof(Id) * 8, "ScalarType id is too large to be stored"); - - auto or_and_advance = [](std::pair result, auto member) -> std::pair { - auto [id, bit_offset] = result; - auto constexpr bits = member_id_field_width(); - return {id | (int64_t(member) & ((uint64_t(1) << bits) - 1)) << bit_offset, bit_offset + bits}; - }; - return reduce_members(or_and_advance, std::pair{}).first; - } - - // create a ScalarType from an id, for c++17 template specialization, - // this is not needed once we migrate to c++20 and can pass literal - // classes as template parameters - static constexpr ScalarType from_id(Id id) { - auto extract_and_advance = [id](auto result, auto member) { - using T = decltype(member); - auto [tuple, bit_offset] = result; - auto constexpr bits = member_id_field_width(); - auto extracted_val = static_cast((int64_t(id) >> bit_offset) & ((uint64_t(1) << bits) - 1)); - auto new_tuple = std::tuple_cat(tuple, std::make_tuple(extracted_val)); - return std::pair{new_tuple, bit_offset + bits}; - }; - - auto [tuple_args, _] = reduce_member_types(extract_and_advance, std::pair, int>{}); - return std::apply([](auto... args) { return ScalarType(args...); }, tuple_args); - } - - constexpr int64_t size_bits() const { - return mantissa + exponent + is_signed(); - } - constexpr bool is_signed() const { - return signed_; - } - constexpr bool is_integer() const { - return exponent == 0; - } - constexpr bool is_floating_point() const { - return exponent > 0; - } - constexpr bool is_ieee_754() const { - return is_floating_point() && finite_values_only == false && nan_repr == NAN_IEEE_754; - } - constexpr bool has_nans() const { - return is_floating_point() && nan_repr != NAN_NONE; - } - constexpr bool has_infs() const { - return is_floating_point() && finite_values_only == false; - } - constexpr bool has_bias() const { - return bias != 0; - } - -#ifndef __CUDACC__ - private: - double _floating_point_max() const { - assert(mantissa <= 52 && exponent <= 11); - - uint64_t max_mantissa = (uint64_t(1) << mantissa) - 1; - if (nan_repr == NAN_EXTD_RANGE_MAX_MIN) { - max_mantissa -= 1; - } - - uint64_t max_exponent = (uint64_t(1) << exponent) - 2; - if (nan_repr == NAN_EXTD_RANGE_MAX_MIN || nan_repr == NAN_NONE) { - assert(exponent < 11); - max_exponent += 1; - } - - // adjust the exponent to match that of a double - // for now we assume the exponent bias is the standard 2^(e-1) -1, (where e - // is the exponent bits), there is some precedent for non-standard biases, - // example `float8_e4m3b11fnuz` here: https://github.com/jax-ml/ml_dtypes - // but to avoid premature over complication we are just assuming the - // standard exponent bias until there is a need to support non-standard - // biases - uint64_t exponent_bias = (uint64_t(1) << (exponent - 1)) - 1; - uint64_t exponent_bias_double = (uint64_t(1) << 10) - 1; // double e = 11 - - uint64_t max_exponent_double = max_exponent - exponent_bias + exponent_bias_double; - - // shift the mantissa into the position for a double and - // the exponent - uint64_t double_raw = (max_mantissa << (52 - mantissa)) | (max_exponent_double << 52); - - return *reinterpret_cast(&double_raw); - } - - constexpr std::variant _raw_max() const { - if (is_floating_point()) { - return {_floating_point_max()}; - } else { - assert(size_bits() < 64 || (size_bits() == 64 && is_signed())); - return {(int64_t(1) << mantissa) - 1}; - } - } - - constexpr std::variant _raw_min() const { - if (is_floating_point()) { - assert(is_signed()); - constexpr uint64_t sign_bit_double = (uint64_t(1) << 63); - - double max = _floating_point_max(); - uint64_t max_raw = *reinterpret_cast(&max); - uint64_t min_raw = max_raw | sign_bit_double; - return {*reinterpret_cast(&min_raw)}; - } else { - assert(!is_signed() || size_bits() <= 64); - if (is_signed()) { - // set the top bit to 1 (i.e. INT64_MIN) and the rest to 0 - // then perform an arithmetic shift right to set all the bits above - // (size_bits() - 1) to 1 - return {INT64_MIN >> (64 - size_bits())}; - } else { - return {int64_t(0)}; - } - } - } - - public: - // Max representable value for this scalar type. - // (accounting for bias if there is one) - constexpr std::variant max() const { - return std::visit([this](auto x) -> std::variant { return {x - bias}; }, _raw_max()); - } - - // Min representable value for this scalar type. - // (accounting for bias if there is one) - constexpr std::variant min() const { - return std::visit([this](auto x) -> std::variant { return {x - bias}; }, _raw_min()); - } -#endif // __CUDACC__ - - public: - std::string str() const { - /* naming generally follows: https://github.com/jax-ml/ml_dtypes - * for floating point types (leading f) the scheme is: - * `float_em[flags]` - * flags: - * - no-flags: means it follows IEEE 754 conventions - * - f: means finite values only (no infinities) - * - n: means nans are supported (non-standard encoding) - * for integer types the scheme is: - * `[u]int[b]` - * - if bias is not present it means its zero - */ - if (is_floating_point()) { - auto ret = - "float" + std::to_string(size_bits()) + "_e" + std::to_string(exponent) + "m" + std::to_string(mantissa); - if (!is_ieee_754()) { - if (finite_values_only) { - ret += "f"; - } - if (nan_repr != NAN_NONE) { - ret += "n"; - } - } - return ret; - } else { - auto ret = ((is_signed()) ? "int" : "uint") + std::to_string(size_bits()); - if (has_bias()) { - ret += "b" + std::to_string(bias); - } - return ret; - } - } - - constexpr bool operator==(ScalarType const& other) const { - return mantissa == other.mantissa && exponent == other.exponent && bias == other.bias && signed_ == other.signed_ && - finite_values_only == other.finite_values_only && nan_repr == other.nan_repr; - } -}; - -using ScalarTypeId = ScalarType::Id; - -// "rust style" names generally following: -// https://github.com/pytorch/pytorch/blob/6d9f74f0af54751311f0dd71f7e5c01a93260ab3/torch/csrc/api/include/torch/types.h#L60-L70 -static inline constexpr auto kS4 = ScalarType::int_(4); -static inline constexpr auto kU4 = ScalarType::uint(4); -static inline constexpr auto kU4B8 = ScalarType::uint(4, 8); -static inline constexpr auto kS8 = ScalarType::int_(8); -static inline constexpr auto kU8 = ScalarType::uint(8); -static inline constexpr auto kU8B128 = ScalarType::uint(8, 128); - -static inline constexpr auto kFE2M1f = ScalarType::float_(2, 1, true, ScalarType::NAN_NONE); -static inline constexpr auto kFE3M2f = ScalarType::float_(3, 2, true, ScalarType::NAN_NONE); -static inline constexpr auto kFE4M3fn = ScalarType::float_(4, 3, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN); -static inline constexpr auto kFE8M0fnu = ScalarType(8, 0, false, 0, true, ScalarType::NAN_EXTD_RANGE_MAX_MIN); -static inline constexpr auto kFE5M2 = ScalarType::float_IEEE754(5, 2); -static inline constexpr auto kFE8M7 = ScalarType::float_IEEE754(8, 7); -static inline constexpr auto kFE5M10 = ScalarType::float_IEEE754(5, 10); - -// Fixed width style names, generally following: -// https://github.com/pytorch/pytorch/blob/6d9f74f0af54751311f0dd71f7e5c01a93260ab3/torch/csrc/api/include/torch/types.h#L47-L57 -static inline constexpr auto kInt4 = kS4; -static inline constexpr auto kUint4 = kU4; -static inline constexpr auto kUint4b8 = kU4B8; -static inline constexpr auto kInt8 = kS8; -static inline constexpr auto kUint8 = kU8; -static inline constexpr auto kUint8b128 = kU8B128; - -static inline constexpr auto kFloat4_e2m1f = kFE2M1f; -static inline constexpr auto kFloat6_e3m2f = kFE3M2f; -static inline constexpr auto kFloat8_e4m3fn = kFE4M3fn; -static inline constexpr auto kFloat8_e5m2 = kFE5M2; -static inline constexpr auto kFloat16_e8m7 = kFE8M7; -static inline constexpr auto kFloat16_e5m10 = kFE5M10; - -// colloquial names -static inline constexpr auto kHalf = kFE5M10; -static inline constexpr auto kFloat16 = kHalf; -static inline constexpr auto kBFloat16 = kFE8M7; - -static inline constexpr auto kFloat16Id = kFloat16.id(); -} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h deleted file mode 100644 index 7c9fd52131..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/source_location.h +++ /dev/null @@ -1,40 +0,0 @@ -/// \file source_location.h -/// \brief Portable `source_location` wrapper. -/// -/// Uses `std::source_location` when available (C++20), otherwise falls -/// back to a minimal stub that returns empty/zero values. - -#pragma once -#include - -/// NOTE: fallback to a minimal source_location implementation -#if defined(__cpp_lib_source_location) -#include - -using source_location_t = std::source_location; - -#else - -struct source_location_fallback { - public: - static constexpr source_location_fallback current() noexcept { - return source_location_fallback{}; - } - constexpr source_location_fallback() noexcept = default; - constexpr unsigned line() const noexcept { - return 0; - } - constexpr unsigned column() const noexcept { - return 0; - } - constexpr const char* file_name() const noexcept { - return ""; - } - constexpr const char* function_name() const noexcept { - return ""; - } -}; - -using source_location_t = source_location_fallback; - -#endif diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h deleted file mode 100644 index 1ae9233a61..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/tensor.h +++ /dev/null @@ -1,605 +0,0 @@ -/// \file tensor.h -/// \brief Tensor validation and symbolic matching utilities. -/// -/// Provides the `TensorMatcher` fluent API for validating tensor shapes, -/// strides, dtypes, and devices at kernel entry points, along with -/// `SymbolicSize`, `SymbolicDType`, and `SymbolicDevice` for capturing -/// and cross-checking tensor metadata across multiple tensors. -/// -/// See the "Tensor Checking" section in the JIT kernel dev guide for -/// usage examples. - -#pragma once -#include - -#include -#include -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#ifdef __CUDACC__ -#include -#elif defined(__HIPCC__) -#include -#endif - -namespace host { - -namespace details { - -inline constexpr auto kAnyDeviceID = -1; -inline constexpr auto kAnySize = static_cast(-1); -inline constexpr auto kNullSize = static_cast(-1); -inline constexpr auto kNullDType = static_cast(18u); -inline constexpr auto kNullDevice = static_cast(-1); - -struct SizeRef; -struct DTypeRef; -struct DeviceRef; - -template -struct _dtype_trait {}; - -template -struct _dtype_trait { - inline static constexpr DLDataType value = { - .code = std::is_signed_v ? DLDataTypeCode::kDLInt : DLDataTypeCode::kDLUInt, - .bits = static_cast(sizeof(T) * 8), - .lanes = 1}; -}; - -template -struct _dtype_trait { - inline static constexpr DLDataType value = { - .code = DLDataTypeCode::kDLFloat, .bits = static_cast(sizeof(T) * 8), .lanes = 1}; -}; - -#ifdef __CUDACC__ -template <> -struct _dtype_trait { - inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat, .bits = 16, .lanes = 1}; -}; -template <> -struct _dtype_trait { - inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1}; -}; -template <> -struct _dtype_trait { - inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat8_e4m3fn, .bits = 8, .lanes = 1}; -}; -#elif defined(__HIPCC__) -template <> -struct _dtype_trait { - inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat, .bits = 16, .lanes = 1}; -}; -template <> -struct _dtype_trait { - inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1}; -}; -#endif - -template -struct _device_trait { - inline static constexpr DLDevice value = {.device_type = Code, .device_id = kAnyDeviceID}; -}; - -template -inline constexpr auto kDTypeList = std::array{_dtype_trait::value...}; - -template -inline constexpr auto kDeviceList = std::array{_device_trait::value...}; - -template -struct PrintAbleSpan { - explicit PrintAbleSpan(std::span data) : data(data) {} - std::span data; -}; - -// define DLDataType comparison and printing in root namespace -inline constexpr auto kDeviceStringMap = [] { - constexpr auto map = std::array, 16>{ - std::pair{DLDeviceType::kDLCPU, "cpu"}, - std::pair{DLDeviceType::kDLCUDA, "cuda"}, - std::pair{DLDeviceType::kDLCUDAHost, "cuda_host"}, - std::pair{DLDeviceType::kDLOpenCL, "opencl"}, - std::pair{DLDeviceType::kDLVulkan, "vulkan"}, - std::pair{DLDeviceType::kDLMetal, "metal"}, - std::pair{DLDeviceType::kDLVPI, "vpi"}, - std::pair{DLDeviceType::kDLROCM, "rocm"}, - std::pair{DLDeviceType::kDLROCMHost, "rocm_host"}, - std::pair{DLDeviceType::kDLExtDev, "ext_dev"}, - std::pair{DLDeviceType::kDLCUDAManaged, "cuda_managed"}, - std::pair{DLDeviceType::kDLOneAPI, "oneapi"}, - std::pair{DLDeviceType::kDLWebGPU, "webgpu"}, - std::pair{DLDeviceType::kDLHexagon, "hexagon"}, - std::pair{DLDeviceType::kDLMAIA, "maia"}, - std::pair{DLDeviceType::kDLTrn, "trn"}, - }; - constexpr auto max_type = stdr::max(map | stdv::keys); - auto result = std::array{}; - for (const auto& [code, name] : map) { - result[static_cast(code)] = name; - } - return result; -}(); - -struct PrintableDevice { - DLDevice device; -}; - -inline auto& operator<<(std::ostream& os, DLDevice device) { - const auto& mapping = kDeviceStringMap; - const auto entry = static_cast(device.device_type); - RuntimeCheck(entry < mapping.size()); - const auto name = mapping[entry]; - RuntimeCheck(!name.empty(), "Unknown device: ", int(device.device_type)); - os << name; - if (device.device_id != kAnyDeviceID && device.device_type != DLDeviceType::kDLCPU) { - os << ":" << device.device_id; - } - return os; -} - -inline auto& operator<<(std::ostream& os, PrintableDevice pd) { - return os << pd.device; -} - -template -inline auto& operator<<(std::ostream& os, PrintAbleSpan span) { - os << "["; - for (const auto i : irange(span.data.size())) { - if (i > 0) { - os << ", "; - } - os << span.data[i]; - } - os << "]"; - return os; -} - -} // namespace details - -/// \brief Check whether `dtype` matches the DLDataType for C++ type `T`. -template -inline bool is_type(DLDataType dtype) { - return dtype == details::_dtype_trait::value; -} - -/** - * \brief A symbolic dimension size that can be bound once and - * verified across multiple tensors. - * - * Create with an optional annotation string for error messages: - * \code - * auto N = SymbolicSize{"num_tokens"}; - * \endcode - * - * Call `verify()` during tensor matching to either bind the first - * observed value or check subsequent values match. Call `unwrap()` - * to retrieve the bound value (panics if unset). - */ -struct SymbolicSize { - public: - SymbolicSize(std::string_view annotation = {}) : m_value(details::kNullSize), m_annotation(annotation) {} - SymbolicSize(const SymbolicSize&) = delete; - SymbolicSize& operator=(const SymbolicSize&) = delete; - - auto get_name() const -> std::string_view { - return m_annotation; - } - - auto set_value(int64_t value) -> void { - RuntimeCheck(!this->has_value(), "Size value already set"); - m_value = value; - } - - auto has_value() const -> bool { - return m_value != details::kNullSize; - } - - auto get_value() const -> std::optional { - return this->has_value() ? std::optional{m_value} : std::nullopt; - } - - auto unwrap(DebugInfo info = {}) const -> int64_t { - RuntimeCheck(info, this->has_value(), "Size value is not set"); - return m_value; - } - - auto verify(int64_t value, const char* prefix, int64_t dim) -> void { - if (this->has_value()) { - if (m_value != value) { - [[unlikely]]; - Panic("Size mismatch for ", m_name_str(prefix, dim), ": expected ", m_value, " but got ", value); - } - } else { - this->set_value(value); - } - } - - auto value_or_name(const char* prefix, int64_t dim) const -> std::string { - if (const auto value = this->get_value()) { - return std::to_string(*value); - } else { - return m_name_str(prefix, dim); - } - } - - private: - auto m_name_str(const char* prefix, int64_t dim) const -> std::string { - std::ostringstream os; - os << prefix << '#' << dim; - if (!m_annotation.empty()) os << "('" << m_annotation << "')"; - return std::move(os).str(); - } - - std::int64_t m_value; - std::string_view m_annotation; -}; - -inline auto operator==(DLDevice lhs, DLDevice rhs) -> bool { - return lhs.device_type == rhs.device_type && lhs.device_id == rhs.device_id; -} - -/** - * \brief A symbolic data type that can be constrained and verified. - * - * Optionally restrict allowed types via `set_options()`. - * Use `verify()` to bind/check the dtype, and `unwrap()` to retrieve it. - */ -struct SymbolicDType { - public: - SymbolicDType() : m_value({details::kNullDType, 0, 0}) {} - SymbolicDType(const SymbolicDType&) = delete; - SymbolicDType& operator=(const SymbolicDType&) = delete; - - auto set_value(DLDataType value) -> void { - RuntimeCheck(!this->has_value(), "Dtype value already set"); - RuntimeCheck( - m_check(value), "Dtype value [", value, "] not in the allowed options: ", details::PrintAbleSpan{m_options}); - m_value = value; - } - - auto has_value() const -> bool { - return m_value.code != details::kNullDType; - } - - auto get_value() const -> std::optional { - return this->has_value() ? std::optional{m_value} : std::nullopt; - } - - auto unwrap(DebugInfo info = {}) const -> DLDataType { - RuntimeCheck(info, this->has_value(), "Dtype value is not set"); - return m_value; - } - - auto set_options(std::span options) -> void { - m_options = options; - } - - template - auto set_options() -> void { - m_options = details::kDTypeList; - } - - auto verify(DLDataType dtype) -> void { - if (this->has_value()) { - RuntimeCheck(m_value == dtype, "DType mismatch: expected ", m_value, " but got ", dtype); - } else { - this->set_value(dtype); - } - } - - template - auto is_type() const -> bool { - return ::host::is_type(m_value); - } - - private: - auto m_check(DLDataType value) const -> bool { - return stdr::empty(m_options) || (stdr::find(m_options, value) != stdr::end(m_options)); - } - - std::span m_options; - DLDataType m_value; -}; - -/** - * \brief A symbolic device that can be constrained and verified. - * - * Optionally restrict allowed device types via - * `set_options()`. The device id can be wildcarded. - */ -struct SymbolicDevice { - public: - SymbolicDevice() : m_value({details::kNullDevice, details::kAnyDeviceID}) {} - SymbolicDevice(const SymbolicDevice&) = delete; - SymbolicDevice& operator=(const SymbolicDevice&) = delete; - - auto set_value(DLDevice value) -> void { - RuntimeCheck(!this->has_value(), "Device value already set"); - RuntimeCheck( - m_check(value), - "Device value [", - details::PrintableDevice{value}, - "] not in the allowed options: ", - details::PrintAbleSpan{m_options}); - m_value = value; - } - - auto has_value() const -> bool { - return m_value.device_type != details::kNullDevice; - } - - auto get_value() const -> std::optional { - return this->has_value() ? std::optional{m_value} : std::nullopt; - } - - auto unwrap(DebugInfo info = {}) const -> DLDevice { - RuntimeCheck(info, this->has_value(), "Device value is not set"); - return m_value; - } - - auto set_options(std::span options) -> void { - m_options = options; - } - - template - auto set_options() -> void { - m_options = details::kDeviceList; - } - - auto verify(DLDevice device) -> void { - if (this->has_value()) { - RuntimeCheck( - m_value == device, - "Device mismatch: expected ", - details::PrintableDevice{m_value}, - " but got ", - details::PrintableDevice{device}); - } else { - this->set_value(device); - } - } - - private: - auto m_check(DLDevice value) const -> bool { - return stdr::empty(m_options) || (stdr::any_of(m_options, [value](const DLDevice& opt) { - // device type must exactly match - if (opt.device_type != value.device_type) return false; - // device id can be wildcarded - return opt.device_id == details::kAnyDeviceID || opt.device_id == value.device_id; - })); - } - - std::span m_options; - DLDevice m_value; -}; - -namespace details { - -template -struct BaseRef { - public: - BaseRef(const BaseRef&) = delete; - BaseRef& operator=(const BaseRef&) = delete; - - auto operator->() const -> T* { - return m_ref; - } - auto operator*() const -> T& { - return *m_ref; - } - auto rebind(T& other) -> void { - m_ref = &other; - } - - explicit BaseRef() : m_ref(&m_cache), m_cache() {} - BaseRef(T& size) : m_ref(&size), m_cache() {} - - private: - T* m_ref; - T m_cache; -}; - -struct SizeRef : BaseRef { - using BaseRef::BaseRef; - SizeRef(int64_t value) { - if (value != kAnySize) { - (**this).set_value(value); - } else { - // otherwise, we can match any size - } - } -}; - -struct DTypeRef : BaseRef { - using BaseRef::BaseRef; - DTypeRef(DLDataType options) { - (**this).set_value(options); - } - DTypeRef(std::initializer_list options) { - (**this).set_options(options); - } - DTypeRef(std::span options) { - (**this).set_options(options); - } -}; - -struct DeviceRef : BaseRef { - using BaseRef::BaseRef; - DeviceRef(DLDevice options) { - (**this).set_value(options); - } - DeviceRef(std::initializer_list options) { - (**this).set_options(options); - } - DeviceRef(std::span options) { - (**this).set_options(options); - } -}; - -} // namespace details - -/** - * \brief Fluent API for validating tensor shape, strides, dtype, and device. - * - * Construct with the expected shape (using `SymbolicSize` or literal - * integers), chain `.with_strides()`, `.with_dtype<...>()`, and - * `.with_device<...>()`, then call `.verify(tensor)`. - * - * Example: - * \code - * auto N = SymbolicSize{"N"}; - * TensorMatcher({N, 128}) - * .with_dtype() - * .with_device() - * .verify(input_tensor); - * \endcode - * - * \note `TensorMatcher` is a move-only temporary. Do not store in a variable. - */ -struct TensorMatcher { - private: - using SizeRef = details::SizeRef; - using DTypeRef = details::DTypeRef; - using DeviceRef = details::DeviceRef; - - public: - TensorMatcher(const TensorMatcher&) = delete; - TensorMatcher& operator=(const TensorMatcher&) = delete; - - explicit TensorMatcher(std::initializer_list shape) : m_shape(shape), m_strides(), m_dtype() {} - - auto with_strides(std::initializer_list strides) && -> TensorMatcher&& { - // no partial update allowed - RuntimeCheck(m_strides.size() == 0, "Strides already specified"); - RuntimeCheck(m_shape.size() == strides.size(), "Strides size must match shape size"); - m_strides = strides; - return std::move(*this); - } - - template - auto with_dtype(DTypeRef&& dtype) && -> TensorMatcher&& { - m_init_dtype(); - m_dtype.rebind(*dtype); - m_dtype->set_options(); - return std::move(*this); - } - - template - auto with_dtype() && -> TensorMatcher&& { - static_assert(sizeof...(Ts) > 0, "At least one dtype option must be specified"); - m_init_dtype(); - m_dtype->set_options(); - return std::move(*this); - } - - template - auto with_device(DeviceRef&& device) && -> TensorMatcher&& { - m_init_device(); - m_device.rebind(*device); - m_device->set_options(); - return std::move(*this); - } - - template - auto with_device() && -> TensorMatcher&& { - static_assert(sizeof...(Codes) > 0, "At least one device option must be specified"); - m_init_device(); - m_device->set_options(); - return std::move(*this); - } - - // once we start verification, we cannot modify anymore - auto verify(tvm::ffi::TensorView view, DebugInfo info = {}) const&& -> const TensorMatcher&& { - try { - m_verify_impl(view); - } catch (PanicError& e) { - auto oss = std::ostringstream{}; - oss << "Tensor match failed for "; - s_print_tensor(oss, view); - oss << " at " << info.file_name() << ":" << info.line() << "\n- Root cause: " << e.root_cause(); - throw PanicError(std::move(oss).str()); - } - return std::move(*this); - } - - private: - static auto s_print_tensor(std::ostringstream& oss, tvm::ffi::TensorView view) -> void { - oss << "Tensor<"; - int64_t dim = 0; - for (const auto& size : view.shape()) { - if (dim++ > 0) oss << ", "; - oss << size; - } - oss << ">[strides=<"; - dim = 0; - for (const auto& stride : view.strides()) { - if (dim++ > 0) { - oss << ", "; - } - oss << stride; - } - oss << ">, dtype=" << view.dtype(); - oss << ", device=" << details::PrintableDevice{view.device()} << "]"; - } - - auto m_verify_impl(tvm::ffi::TensorView view) const -> void { - const auto dim = static_cast(view.dim()); - RuntimeCheck(dim == m_shape.size(), "Tensor dimension mismatch: expected ", m_shape.size(), " but got ", dim); - for (const auto i : irange(dim)) { - m_shape[i]->verify(view.size(i), "shape", i); - } - if (m_has_strides()) { - for (const auto i : irange(dim)) { - if (view.size(i) != 1 || !m_strides[i]->has_value()) { - // skip stride check for size 1 dimension - m_strides[i]->verify(view.stride(i), "stride", i); - } - } - } else { - RuntimeCheck(view.is_contiguous(), "Tensor is not contiguous as expected"); - } - // since we may double verify, we will force to check - m_dtype->verify(view.dtype()); - m_device->verify(view.device()); - } - - auto m_init_dtype() -> void { - RuntimeCheck(!m_has_dtype, "DType already specified"); - m_has_dtype = true; - } - - auto m_init_device() -> void { - RuntimeCheck(!m_has_device, "Device already specified"); - m_has_device = true; - } - - auto m_has_strides() const -> bool { - return !m_strides.empty(); - } - - std::span m_shape; - std::span m_strides; - DTypeRef m_dtype; - DeviceRef m_device; - bool m_has_dtype = false; - bool m_has_device = false; -}; - -} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh deleted file mode 100644 index 1adc821706..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/tile.cuh +++ /dev/null @@ -1,62 +0,0 @@ -/// \file tile.cuh -/// \brief Tiled memory access helpers for coalesced global memory I/O. -/// -/// `tile::Memory` represents a contiguous memory region where multiple -/// threads cooperatively load/store elements. The three factory methods -/// determine the thread group: -/// - `thread()` - single thread (no tiling). -/// - `warp()` - all threads in a warp cooperate. -/// - `cta()` - all threads in the CTA cooperate. - -#pragma once -#include - -#include - -namespace device::tile { - -/** - * \brief Represents a contiguous memory region for cooperative tiled access. - * - * Each instance is parameterized by an element type `T` and bound to a - * specific thread id (`tid`) within a group of `tsize` threads. - * - * \tparam T The storage element type (e.g. `AlignedVector, 4>`). - */ -template -struct Memory { - public: - SGL_DEVICE constexpr Memory(uint32_t tid, uint32_t tsize) : tid(tid), tsize(tsize) {} - /// \brief Create a Memory accessor for a single thread (no cooperation). - SGL_DEVICE static constexpr Memory thread() { - return Memory{0, 1}; - } - /// \brief Create a Memory accessor distributed across warp threads. - SGL_DEVICE static Memory warp(int warp_threads = kWarpThreads) { - return Memory{static_cast(threadIdx.x % warp_threads), static_cast(warp_threads)}; - } - /// \brief Create a Memory accessor distributed across all CTA threads. - SGL_DEVICE static Memory cta(int cta_threads = blockDim.x) { - return Memory{static_cast(threadIdx.x), static_cast(cta_threads)}; - } - /// \brief Load one element from `ptr` at the position assigned to this thread. - /// \param ptr Base pointer (cast to `const T*`). - /// \param offset Optional tile offset (multiplied by `tsize`). - SGL_DEVICE T load(const void* ptr, int64_t offset = 0) const { - return static_cast(ptr)[tid + offset * tsize]; - } - /// \brief Store one element to `ptr` at the position assigned to this thread. - SGL_DEVICE void store(void* ptr, T val, int64_t offset = 0) const { - static_cast(ptr)[tid + offset * tsize] = val; - } - /// \brief Check whether this thread's element index is within bounds. - SGL_DEVICE bool in_bound(int64_t element_count, int64_t offset = 0) const { - return tid + offset * tsize < element_count; - } - - private: - uint32_t tid; - uint32_t tsize; -}; - -} // namespace device::tile diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh deleted file mode 100644 index a7a5346196..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/type.cuh +++ /dev/null @@ -1,120 +0,0 @@ -/// \file type.cuh -/// \brief Dtype trait system for CUDA scalar/packed types. -/// -/// `dtype_trait` provides per-type metadata: packed type alias, -/// conversion functions (`from`), and unary/binary math operations. -/// Use `device::cast(from_value)` for type conversion on device. -/// -/// Registered types: -/// | Scalar | Packed (x2) | Notes | -/// |-----------|-------------|-------------------------------| -/// | `fp32_t` | `fp32x2_t` | Full math ops (abs,sqrt,...) | -/// | `fp16_t` | `fp16x2_t` | Conversion only | -/// | `bf16_t` | `bf16x2_t` | Conversion only | -/// | `fp32x2_t`| `fp32x4_t` | Packed float2 <-> half2/bf162 | - -#pragma once -#include - -template -struct dtype_trait {}; - -#define SGL_REGISTER_DTYPE_TRAIT(TYPE, PACK2, ...) \ - template <> \ - struct dtype_trait { \ - using self_t = TYPE; \ - using packed_t = PACK2; \ - template \ - SGL_DEVICE static self_t from(const S& value) { \ - return static_cast(value); \ - } \ - __VA_ARGS__ \ - } - -#define SGL_REGISTER_TYPE_END static_assert(true) - -#define SGL_REGISTER_FROM_FUNCTION(FROM, FN) \ - SGL_DEVICE static self_t from(const FROM& x) { \ - return FN(x); \ - } \ - static_assert(true) - -#define SGL_REGISTER_UNARY_FUNCTION(NAME, FN) \ - SGL_DEVICE static self_t NAME(const self_t& x) { \ - return FN(x); \ - } \ - static_assert(true) - -#define SGL_REGISTER_BINARY_FUNCTION(NAME, FN) \ - SGL_DEVICE static self_t NAME(const self_t& x, const self_t& y) { \ - return FN(x, y); \ - } \ - static_assert(true) - -SGL_REGISTER_DTYPE_TRAIT( - fp32_t, fp32x2_t, SGL_REGISTER_TYPE_END; // - SGL_REGISTER_FROM_FUNCTION(fp16_t, __half2float); - SGL_REGISTER_FROM_FUNCTION(bf16_t, __bfloat162float); - SGL_REGISTER_UNARY_FUNCTION(abs, fabsf); - SGL_REGISTER_UNARY_FUNCTION(sqrt, sqrtf); - SGL_REGISTER_UNARY_FUNCTION(rsqrt, rsqrtf); - SGL_REGISTER_UNARY_FUNCTION(exp, expf); - SGL_REGISTER_UNARY_FUNCTION(sin, sinf); - SGL_REGISTER_UNARY_FUNCTION(cos, cosf); - SGL_REGISTER_BINARY_FUNCTION(max, fmaxf); - SGL_REGISTER_BINARY_FUNCTION(min, fminf);); -SGL_REGISTER_DTYPE_TRAIT(fp16_t, fp16x2_t); -SGL_REGISTER_DTYPE_TRAIT(bf16_t, bf16x2_t); - -/// TODO: Add ROCM implementation -SGL_REGISTER_DTYPE_TRAIT( - fp32x2_t, fp32x4_t, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp16x2_t, __half22float2); - SGL_REGISTER_FROM_FUNCTION(bf16x2_t, __bfloat1622float2);); - -SGL_REGISTER_DTYPE_TRAIT( - fp16x2_t, void, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp32x2_t, __float22half2_rn);); - -SGL_REGISTER_DTYPE_TRAIT( - bf16x2_t, void, SGL_REGISTER_TYPE_END; SGL_REGISTER_FROM_FUNCTION(fp32x2_t, __float22bfloat162_rn);); - -#ifndef USE_ROCM -SGL_REGISTER_DTYPE_TRAIT(fp8_e4m3_t, fp8x2_e4m3_t); -#endif - -#undef SGL_REGISTER_DTYPE_TRAIT -#undef SGL_REGISTER_FROM_FUNCTION - -/// \brief Alias: the packed (x2) type for `T`. -template -using packed_t = typename dtype_trait::packed_t; - -namespace device { - -/** - * \brief Cast a value from type `From` to type `To` on device. - * - * Dispatches through `dtype_trait::from()`, which uses the appropriate - * CUDA intrinsic (e.g. `__half2float`, `__float22half2_rn`). - */ -template -SGL_DEVICE To cast(const From& value) { - return dtype_trait::from(value); -} - -} // namespace device - -// --------------------------------------------------------------------------- -// FP8 max clamp value — platform-dependent -// CUDA (e4m3fn): 448.0f -// AMD FNUZ (e4m3fnuz): 224.0f -// AMD E4M3 (e4m3fn): 448.0f -// --------------------------------------------------------------------------- -#ifndef USE_ROCM -constexpr float kFP8E4M3Max = 448.0f; -#else // USE_ROCM -#if HIP_FP8_TYPE_FNUZ -constexpr float kFP8E4M3Max = 224.0f; -#else // HIP_FP8_TYPE_E4M3 -constexpr float kFP8E4M3Max = 448.0f; -#endif // HIP_FP8_TYPE_FNUZ -#endif // USE_ROCM diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh deleted file mode 100644 index 2dd6f3dc93..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.cuh +++ /dev/null @@ -1,333 +0,0 @@ -/// \file utils.cuh -/// \brief Core CUDA/device utilities: type aliases, PDL helpers, -/// typed pointer access, kernel launch wrapper, and error checking. -/// -/// This header is included (directly or transitively) by nearly every -/// JIT kernel. It provides: -/// - Scalar/packed type aliases (`fp16_t`, `bf16_t`, `fp8_e4m3_t`, ...). -/// - `SGL_DEVICE` macro (forced-inline device function qualifier). -/// - `kWarpThreads` constant (32). -/// - PDL (Programmatic Dependent Launch) helpers for Hopper (sm_90+). -/// - Typed `load_as` / `store_as` for void-pointer access. -/// - `pointer::offset` for safe void-pointer arithmetic. -/// - `host::LaunchKernel` - kernel launcher with optional PDL. -/// - `host::RuntimeDeviceCheck` - CUDA error checking. - -#pragma once - -#include - -#include -#include - -#include -#include -#include -#ifndef USE_ROCM -#include -#include -#include -#include -#else -#include -#include -#include -#ifndef __grid_constant__ -#define __grid_constant__ -#endif -using cudaError_t = hipError_t; -using cudaStream_t = hipStream_t; -using cudaLaunchConfig_t = hipLaunchConfig_t; -using cudaLaunchAttribute = hipLaunchAttribute; -inline constexpr auto cudaSuccess = hipSuccess; -#define cudaStreamPerThread hipStreamPerThread -#define cudaGetErrorString hipGetErrorString -#define cudaGetLastError hipGetLastError -#define cudaLaunchKernel hipLaunchKernel -#define cudaMemcpyAsync hipMemcpyAsync -#define cudaMemcpyHostToDevice hipMemcpyHostToDevice -#define cudaMemcpyDeviceToHost hipMemcpyDeviceToHost -#endif - -#ifndef USE_ROCM -using fp32_t = float; -using fp16_t = __half; -using bf16_t = __nv_bfloat16; -using fp8_e4m3_t = __nv_fp8_e4m3; -using fp8_e5m2_t = __nv_fp8_e5m2; - -using fp32x2_t = float2; -using fp16x2_t = __half2; -using bf16x2_t = __nv_bfloat162; -using fp8x2_e4m3_t = __nv_fp8x2_e4m3; -using fp8x2_e5m2_t = __nv_fp8x2_e5m2; - -using fp32x4_t = float4; -#else -using fp32_t = float; -using fp16_t = __half; -using bf16_t = __hip_bfloat16; -using fp8_e4m3_t = uint8_t; -using fp8_e5m2_t = uint8_t; -using fp32x2_t = float2; -using fp16x2_t = half2; -using bf16x2_t = __hip_bfloat162; -using fp8x2_e4m3_t = uint16_t; -using fp8x2_e5m2_t = uint16_t; -using fp32x4_t = float4; -#endif - -/* - * LDG Support - */ -#ifndef USE_ROCM -#define SGLANG_LDG(arg) __ldg(arg) -#else -#define SGLANG_LDG(arg) *(arg) -#endif - -// DLPack device type for the current platform -#ifndef USE_ROCM -inline constexpr auto kDLGPU = kDLCUDA; -#else -inline constexpr auto kDLGPU = kDLROCM; -#endif - -namespace device { - -/// \brief Macro: forced-inline device function qualifier. -#define SGL_DEVICE __forceinline__ __device__ - -// Architecture detection: SGL_CUDA_ARCH is injected by load_jit() and is -// available in both host and device compilation passes, whereas __CUDA_ARCH__ -// is only defined by nvcc during the device pass. -#if !defined(USE_ROCM) -#if !defined(SGL_CUDA_ARCH) -#error "SGL_CUDA_ARCH is not defined. JIT compilation must inject -DSGL_CUDA_ARCH via load_jit()." -#endif -#if defined(__CUDA_ARCH__) -static_assert( - __CUDA_ARCH__ == SGL_CUDA_ARCH, "SGL_CUDA_ARCH mismatch: injected arch flag does not match device target"); -#endif -#define SGL_ARCH_HOPPER_OR_GREATER (SGL_CUDA_ARCH >= 900) -#define SGL_ARCH_BLACKWELL_OR_GREATER ((SGL_CUDA_ARCH >= 1000) && (CUDA_VERSION >= 12090)) -#else // USE_ROCM -#define SGL_ARCH_HOPPER_OR_GREATER 0 -#define SGL_ARCH_BLACKWELL_OR_GREATER 0 -#endif - -// Maximum vector size in bytes supported by current architecture. -// Pre-Blackwell / AMD: 128-bit (16 bytes) -// Blackwell or greater: 256-bit (32 bytes) -inline constexpr std::size_t kMaxVecBytes = SGL_ARCH_BLACKWELL_OR_GREATER ? 32 : 16; - -/// \brief Number of threads per warp (always 32 on NVIDIA/AMD GPUs). -inline constexpr auto kWarpThreads = 32u; -/// \brief Full warp active mask (all 32 lanes). -#ifndef USE_ROCM -inline constexpr auto kFullMask = 0xffffffffu; -#else -inline constexpr auto kFullMask = 0xffffffffffffffffULL; -#endif - -/** - * \brief PDL (Programmatic Dependent Launch): wait for the primary kernel. - * - * On Hopper (sm_90+), inserts a `griddepcontrol.wait` instruction to - * synchronize with a preceding kernel in the same stream. On older - * architectures or ROCm this is a no-op. - */ -template -SGL_DEVICE void PDLWaitPrimary() { -#if SGL_ARCH_HOPPER_OR_GREATER - if constexpr (kUsePDL) { - asm volatile("griddepcontrol.wait;" ::: "memory"); - } -#endif -} - -/** - * \brief PDL: trigger dependent (secondary) kernel launch. - * - * On Hopper (sm_90+), inserts a `griddepcontrol.launch_dependents` - * instruction. On older architectures or ROCm this is a no-op. - */ -template -SGL_DEVICE void PDLTriggerSecondary() { -#if SGL_ARCH_HOPPER_OR_GREATER - if constexpr (kUsePDL) { - asm volatile("griddepcontrol.launch_dependents;" :::); - } -#endif -} - -template -SGL_DEVICE constexpr auto div_ceil(T a, U b) { - return (a + b - 1) / b; -} - -/** - * \brief Load data with the specified type and offset from a void pointer. - * \tparam T The type to load. - * \param ptr The base pointer. - * \param offset The offset in number of elements of type T. - */ -template -SGL_DEVICE T load_as(const void* ptr, int64_t offset = 0) { - return static_cast(ptr)[offset]; -} - -/** - * \brief Store data with the specified type and offset to a void pointer. - * \tparam T The type to store. - * \param ptr The base pointer. - * \param val The value to store. - * \param offset The offset in number of elements of type T. - * \note we use type_identity_t to force the caller to explicitly specify - * the template parameter `T`, which can avoid accidentally using the wrong type. - */ -template -SGL_DEVICE void store_as(void* ptr, std::type_identity_t val, int64_t offset = 0) { - static_cast(ptr)[offset] = val; -} - -/// \brief Safe void-pointer arithmetic (byte-level by default). -namespace pointer { - -// we only allow void * pointer arithmetic for safety - -template -SGL_DEVICE auto offset(void* ptr, U... offset) -> void* { - return static_cast(ptr) + (... + offset); -} - -template -SGL_DEVICE auto offset(const void* ptr, U... offset) -> const void* { - return static_cast(ptr) + (... + offset); -} - -} // namespace pointer - -} // namespace device - -namespace host { - -/** - * \brief Check the CUDA error code and panic with location info on failure. - */ -inline void RuntimeDeviceCheck(::cudaError_t error, DebugInfo location = {}) { - if (error != ::cudaSuccess) { - [[unlikely]]; - ::host::panic(location, "CUDA error: ", ::cudaGetErrorString(error)); - } -} - -/// \brief Check the last CUDA error (calls `cudaGetLastError`). -inline void RuntimeDeviceCheck(DebugInfo location = {}) { - return RuntimeDeviceCheck(::cudaGetLastError(), location); -} - -/** - * \brief Kernel launcher with automatic stream resolution and PDL support. - * - * Usage: - * \code - * host::LaunchKernel(grid, block, device) - * .enable_pdl(true) - * (my_kernel, arg1, arg2); - * \endcode - * - * The constructor resolves the CUDA stream from a `DLDevice` (via - * `TVMFFIEnvGetStream`) or accepts a raw `cudaStream_t`. The call - * operator launches the kernel and checks for errors. - */ -struct LaunchKernel { - public: - explicit LaunchKernel( - dim3 grid_dim, - dim3 block_dim, - DLDevice device, - std::size_t dynamic_shared_mem_bytes = 0, - DebugInfo location = {}) noexcept - : m_config(s_make_config(grid_dim, block_dim, resolve_device(device), dynamic_shared_mem_bytes)), - m_location(location) {} - - explicit LaunchKernel( - dim3 grid_dim, - dim3 block_dim, - cudaStream_t stream, - std::size_t dynamic_shared_mem_bytes = 0, - DebugInfo location = {}) noexcept - : m_config(s_make_config(grid_dim, block_dim, stream, dynamic_shared_mem_bytes)), m_location(location) {} - - LaunchKernel(const LaunchKernel&) = delete; - LaunchKernel& operator=(const LaunchKernel&) = delete; - - static auto resolve_device(DLDevice device) -> cudaStream_t { - return static_cast(::TVMFFIEnvGetStream(device.device_type, device.device_id)); - } - - auto enable_pdl(bool enabled = true) -> LaunchKernel& { -#ifdef USE_ROCM - (void)enabled; - m_config.numAttrs = 0; -#else - if (enabled) { - auto& attr = m_attrs[m_config.numAttrs++]; - attr.id = cudaLaunchAttributeProgrammaticStreamSerialization; - attr.val.programmaticStreamSerializationAllowed = true; - m_config.attrs = m_attrs; - } -#endif - return *this; - } - - auto enable_cluster(dim3 cluster_dim) -> LaunchKernel& { -#ifdef USE_ROCM - (void)cluster_dim; -#else - auto& attr = m_attrs[m_config.numAttrs++]; - attr.id = cudaLaunchAttributeClusterDimension; - attr.val.clusterDim = {cluster_dim.x, cluster_dim.y, cluster_dim.z}; - m_config.attrs = m_attrs; -#endif - return *this; - } - - template - auto operator()(T&& kernel, Args&&... args) const -> void { -#ifdef USE_ROCM - hipLaunchKernelGGL( - std::forward(kernel), - m_config.gridDim, - m_config.blockDim, - m_config.dynamicSmemBytes, - m_config.stream, - std::forward(args)...); - RuntimeDeviceCheck(m_location); -#else - RuntimeDeviceCheck(::cudaLaunchKernelEx(&m_config, kernel, std::forward(args)...), m_location); -#endif - } - - private: - static auto s_make_config( // Make a config for kernel launch - dim3 grid_dim, - dim3 block_dim, - cudaStream_t stream, - std::size_t smem) -> cudaLaunchConfig_t { - auto config = ::cudaLaunchConfig_t{}; - config.gridDim = grid_dim; - config.blockDim = block_dim; - config.dynamicSmemBytes = smem; - config.stream = stream; - config.numAttrs = 0; - return config; - } - - cudaLaunchConfig_t m_config; - const DebugInfo m_location; - cudaLaunchAttribute m_attrs[2]; -}; - -} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h b/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h deleted file mode 100644 index 3226f79ddc..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/utils.h +++ /dev/null @@ -1,186 +0,0 @@ -/// \file utils.h -/// \brief Host-side C++ utilities used by JIT kernel wrappers. -/// -/// Provides: -/// - `DebugInfo` - wraps `std::source_location` for error reporting. -/// - `RuntimeCheck` - runtime assertion with formatted error messages. -/// - `Panic` - unconditional abort with formatted error messages. -/// - `pointer::offset` - safe void-pointer arithmetic (host side). -/// - `div_ceil` - integer ceiling division. -/// - `dtype_bytes` - byte width of a `DLDataType`. -/// - `irange` - Python-style integer range for range-for loops. - -#pragma once - -// ref: https://forums.developer.nvidia.com/t/c-20s-source-location-compilation-error-when-using-nvcc-12-1/258026/3 -#ifdef __CUDACC__ -#include -#if CUDA_VERSION <= 12010 - -#pragma push_macro("__cpp_consteval") -#pragma push_macro("_NODISCARD") -#pragma push_macro("__builtin_LINE") - -#pragma clang diagnostic push -#pragma clang diagnostic ignored "-Wbuiltin-macro-redefined" -#define __cpp_consteval 201811L -#pragma clang diagnostic pop - -#ifdef _NODISCARD -#undef _NODISCARD -#define _NODISCARD -#endif - -#define consteval constexpr - -#include "source_location.h" - -#undef consteval -#pragma pop_macro("__cpp_consteval") -#pragma pop_macro("_NODISCARD") -#else // __CUDACC__ && CUDA_VERSION > 12010 -#include "source_location.h" -#endif -#else // no __CUDACC__ -#include "source_location.h" -#endif - -#include - -#include -#include -#include -#include -#include -#include - -namespace host { - -template -inline constexpr bool dependent_false_v = false; - -/// \brief Source-location wrapper for debug/error messages. -struct DebugInfo : public source_location_t { - DebugInfo(source_location_t loc = source_location_t::current()) : source_location_t(loc) {} -}; - -/// \brief Exception type thrown by `RuntimeCheck` and `Panic`. -struct PanicError : public std::runtime_error { - public: - explicit PanicError(std::string msg) : runtime_error(msg), m_message(std::move(msg)) {} - auto root_cause() const -> std::string_view { - const auto str = std::string_view{m_message}; - const auto pos = str.find(": "); - return pos == std::string_view::npos ? str : str.substr(pos + 2); - } - - private: - std::string m_message; -}; - -/// \brief Unconditionally abort with a formatted error message. -template -[[noreturn]] -inline auto panic(DebugInfo location, Args&&... args) -> void { - std::ostringstream os; - os << "Runtime check failed at " << location.file_name() << ":" << location.line(); - if constexpr (sizeof...(args) > 0) { - os << ": "; - (os << ... << std::forward(args)); - } else { - os << " in " << location.function_name(); - } - throw PanicError(std::move(os).str()); -} - -/** - * \brief Runtime assertion: panics with a formatted message when `condition` - * is false. Extra `args` are streamed to the error message. - * - * Example: - * \code - * RuntimeCheck(n > 0, "n must be positive, got ", n); - * \endcode - */ -template -struct RuntimeCheck { - template - explicit RuntimeCheck(Cond&& condition, Args&&... args, DebugInfo location = {}) { - if (condition) return; - [[unlikely]] ::host::panic(location, std::forward(args)...); - } - template - explicit RuntimeCheck(DebugInfo location, Cond&& condition, Args&&... args) { - if (condition) return; - [[unlikely]] ::host::panic(location, std::forward(args)...); - } -}; - -template -struct Panic { - explicit Panic(Args&&... args, DebugInfo location = {}) { - ::host::panic(location, std::forward(args)...); - } - explicit Panic(DebugInfo location, Args&&... args) { - ::host::panic(location, std::forward(args)...); - } - [[noreturn]] ~Panic() { - std::terminate(); - } -}; - -template -explicit RuntimeCheck(Cond&&, Args&&...) -> RuntimeCheck; - -template -explicit RuntimeCheck(DebugInfo, Cond&&, Args&&...) -> RuntimeCheck; - -template -explicit Panic(Args&&...) -> Panic; - -template -explicit Panic(DebugInfo, Args&&...) -> Panic; - -namespace pointer { - -// we only allow void * pointer arithmetic for safety - -template -inline auto offset(void* ptr, U... offset) -> void* { - return static_cast(ptr) + (... + offset); -} - -template -inline auto offset(const void* ptr, U... offset) -> const void* { - return static_cast(ptr) + (... + offset); -} - -} // namespace pointer - -/// \brief Integer ceiling division: ceil(a / b). -template -inline constexpr auto div_ceil(T a, U b) { - return (a + b - 1) / b; -} - -/// \brief Returns the byte width of a DLPack data type. -inline auto dtype_bytes(DLDataType dtype) -> std::size_t { - return static_cast(dtype.bits / 8); -} - -namespace stdr = std::ranges; -namespace stdv = stdr::views; - -/// \brief Python-style integer range: `irange(n)` -> `[0, n)`. -template -inline auto irange(T end) { - return stdv::iota(static_cast(0), end); -} - -/// \brief Python-style integer range: `irange(start, end)` -> `[start, end)`. -template -inline auto irange(T start, T end) { - return stdv::iota(start, end); -} - -} // namespace host diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh deleted file mode 100644 index 67f388679f..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/vec.cuh +++ /dev/null @@ -1,118 +0,0 @@ -/// \file vec.cuh -/// \brief Aligned vector types for coalesced global memory access. -/// -/// `AlignedVector` wraps `N` elements of type `T` in a naturally -/// aligned struct so that the compiler emits wide (vectorized) load/store -/// instructions (e.g. `LDG.128`). The maximum supported vector width is -/// 256 bits (32 bytes), matching CUDA's widest vector load. - -#pragma once -#include - -#include -#include - -namespace device { - -namespace details { - -/// \brief Maps byte-width to the corresponding unsigned integer type. -template -struct uint_trait {}; - -template <> -struct uint_trait<1> { - using type = uint8_t; -}; - -template <> -struct uint_trait<2> { - using type = uint16_t; -}; - -template <> -struct uint_trait<4> { - using type = uint32_t; -}; - -template <> -struct uint_trait<8> { - using type = uint64_t; -}; - -/// \brief Alias: maps `sizeof(T)` to matching unsigned int type. -template -using sized_int = typename uint_trait::type; - -} // namespace details - -/// \brief Raw aligned storage for `N` elements of type `T`. -template -struct alignas(sizeof(T) * N) AlignedStorage { - T data[N]; -}; - -/** - * \brief Aligned vector for vectorized memory access on GPU. - * - * Stores `N` elements of type `T` with natural alignment so that a single - * `load`/`store` call compiles to a wide memory transaction. - * - * \tparam T Element type (e.g. `fp16_t`, `bf16_t`, `float`). - * \tparam N Number of elements. Must be a power of two and - * `sizeof(T) * N <= 32` (256 bits). - * - * Example: - * \code - * AlignedVector vec; // 16 bytes, 128-bit aligned - * vec.load(input_ptr, tid); // vectorized load - * vec[0] = vec[0] + 1; - * vec.store(output_ptr, tid); // vectorized store - * \endcode - */ -template -struct AlignedVector { - private: - static_assert( - (N > 0 && (N & (N - 1)) == 0) && sizeof(T) * N <= kMaxVecBytes, - "CUDA vector size exceeds arch limit: max 16 bytes on pre-Blackwell/AMD, " - "32 bytes on Blackwell or greater"); - using element_t = typename details::sized_int; - using storage_t = AlignedStorage; - - public: - /// \brief Vectorized load from `ptr` at the given element `offset`. - SGL_DEVICE void load(const void* ptr, int64_t offset = 0) { - m_storage = reinterpret_cast(ptr)[offset]; - } - /// \brief Vectorized store to `ptr` at the given element `offset`. - SGL_DEVICE void store(void* ptr, int64_t offset = 0) const { - reinterpret_cast(ptr)[offset] = m_storage; - } - /// \brief Fill all N elements with the same `value`. - SGL_DEVICE void fill(T value) { - const auto store_value = *reinterpret_cast(&value); -#pragma unroll - for (std::size_t i = 0; i < N; ++i) { - m_storage.data[i] = store_value; - } - } - - SGL_DEVICE auto operator[](std::size_t idx) -> T& { - return reinterpret_cast(&m_storage)[idx]; - } - SGL_DEVICE auto operator[](std::size_t idx) const -> T { - return reinterpret_cast(&m_storage)[idx]; - } - SGL_DEVICE auto data() -> T* { - return reinterpret_cast(&m_storage); - } - SGL_DEVICE auto data() const -> const T* { - return reinterpret_cast(&m_storage); - } - - private: - storage_t m_storage; -}; - -} // namespace device diff --git a/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh b/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh deleted file mode 100644 index 9d82efae1e..0000000000 --- a/lightllm/third_party/sglang_jit/include/sgl_kernel/warp.cuh +++ /dev/null @@ -1,56 +0,0 @@ -/// \file warp.cuh -/// \brief Warp-level reduction primitives. - -#pragma once -#include -#include - -namespace device::warp { - -/// \brief Full warp active mask. -#ifndef USE_ROCM -static constexpr uint32_t kFullMask = 0xffffffffu; -using mask_t = uint32_t; -#else -static constexpr uint64_t kFullMask = 0xffffffffffffffffULL; -using mask_t = uint64_t; -#endif - -/** - * \brief Warp-level sum reduction. - * - * On CUDA: uses __shfl_xor_sync with width=32. - * On HIP: uses __shfl_xor with explicit width parameter (supports wave64 sub-groups). - */ -template -SGL_DEVICE T reduce_sum(T value, mask_t active_mask = kFullMask) { - static_assert(kNumThreads >= 1 && kNumThreads <= kWarpThreads); - static_assert(std::has_single_bit(kNumThreads), "must be pow of 2"); -#pragma unroll - for (int mask = kNumThreads / 2; mask > 0; mask >>= 1) -#ifndef USE_ROCM - value = value + __shfl_xor_sync(active_mask, value, mask, 32); -#else - value = value + __shfl_xor(value, mask, kNumThreads); -#endif - return value; -} - -/** - * \brief Warp-level max reduction. - */ -template -SGL_DEVICE T reduce_max(T value, mask_t active_mask = kFullMask) { - static_assert(kNumThreads >= 1 && kNumThreads <= kWarpThreads); - static_assert(std::has_single_bit(kNumThreads), "must be pow of 2"); -#pragma unroll - for (int mask = kNumThreads / 2; mask > 0; mask >>= 1) -#ifndef USE_ROCM - value = math::max(value, __shfl_xor_sync(active_mask, value, mask, 32)); -#else - value = math::max(value, __shfl_xor(value, mask, kNumThreads)); -#endif - return value; -} - -} // namespace device::warp diff --git a/lightllm/third_party/sglang_jit/jit_utils.py b/lightllm/third_party/sglang_jit/jit_utils.py deleted file mode 100644 index 4096c16bb4..0000000000 --- a/lightllm/third_party/sglang_jit/jit_utils.py +++ /dev/null @@ -1,432 +0,0 @@ -from __future__ import annotations - -import functools -import importlib.util -import logging -import os -import pathlib -from contextlib import contextmanager -from dataclasses import dataclass -from typing import ( - TYPE_CHECKING, - Any, - Callable, - Dict, - List, - Optional, - Tuple, - TypeAlias, - TypeVar, - Union, -) - -import torch - -if TYPE_CHECKING: - from tvm_ffi import Module - -F = TypeVar("F", bound=Callable[..., Any]) -_FULL_TEST_ENV_VAR = "SGLANG_JIT_KERNEL_RUN_FULL_TESTS" - -logger = logging.getLogger(__name__) - - -def is_in_ci() -> bool: - return os.getenv("SGLANG_IS_IN_CI", "").lower() in ("1", "true", "yes", "y") - - -def should_run_full_tests() -> bool: - return os.getenv(_FULL_TEST_ENV_VAR, "false").lower() == "true" - - -def get_ci_test_range(full_range: List[Any], ci_range: List[Any]) -> List[Any]: - if should_run_full_tests(): - return full_range - return ci_range if is_in_ci() else full_range - - -def cache_once(fn: F) -> F: - """ - NOTE: `functools.lru_cache` is not compatible with `torch.compile` - So we manually implement a simple cache_once decorator to replace it. - """ - result_map = {} - - @functools.wraps(fn) - def wrapper(*args, **kwargs): - key = (args, tuple(sorted(kwargs.items()))) - if key not in result_map: - result_map[key] = fn(*args, **kwargs) - return result_map[key] - - return wrapper # type: ignore - - -def _make_wrapper(tup: Tuple[str, str]) -> str: - export_name, kernel_name = tup - return f"TVM_FFI_DLL_EXPORT_TYPED_FUNC({export_name}, ({kernel_name}));" - - -@cache_once -def _resolve_kernel_path() -> pathlib.Path: - cur_dir = pathlib.Path(__file__).parent.resolve() - - # first, try this directory structure - def _environment_install(): - candidate = cur_dir.resolve() - if (candidate / "include").exists() and (candidate / "csrc").exists(): - return candidate - return None - - def _package_install(): - # TODO: support find path by package - return None - - path = _environment_install() or _package_install() - if path is None: - raise RuntimeError("Cannot find sglang.jit_kernel path") - return path - - -KERNEL_PATH = _resolve_kernel_path() -DEFAULT_INCLUDE = [str(KERNEL_PATH / "include")] -DEFAULT_CFLAGS = ["-std=c++20", "-O3"] -DEFAULT_LDFLAGS = [] -CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, str, bool, torch.dtype] - - -class CPPArgList(list[str]): - def __str__(self) -> str: - return ", ".join(self) - - -CPP_DTYPE_MAP = { - torch.float: "fp32_t", - torch.float16: "fp16_t", - torch.float8_e4m3fn: "fp8_e4m3_t", - torch.bfloat16: "bf16_t", - torch.int8: "int8_t", - torch.int32: "int32_t", - torch.int64: "int64_t", -} - - -# AMD/ROCm note: -@cache_once -def is_hip_runtime() -> bool: - return bool(torch.version.hip) - - -# MThreads/MUSA note: -@cache_once -def is_musa_runtime() -> bool: - return hasattr(torch.version, "musa") and torch.version.musa is not None - - -def make_cpp_args(*args: CPP_TEMPLATE_TYPE) -> CPPArgList: - def _convert(arg: CPP_TEMPLATE_TYPE) -> str: - if isinstance(arg, bool): - return "true" if arg else "false" - if isinstance(arg, (int, str, float)): - return str(arg) - if isinstance(arg, torch.dtype): - return CPP_DTYPE_MAP[arg] - raise TypeError(f"Unsupported argument type for cpp template: {type(arg)}") - - return CPPArgList(_convert(arg) for arg in args) - - -def load_jit( - *args: str, - cpp_files: List[str] | None = None, - cuda_files: List[str] | None = None, - cpp_wrappers: List[Tuple[str, str]] | None = None, - cuda_wrappers: List[Tuple[str, str]] | None = None, - extra_cflags: List[str] | None = None, - extra_cuda_cflags: List[str] | None = None, - extra_ldflags: List[str] | None = None, - extra_include_paths: List[str] | None = None, - extra_dependencies: List[str] | None = None, - build_directory: str | None = None, - header_only: bool = True, -) -> Module: - """ - Loading a JIT module from C++/CUDA source files. - We define a wrapper as a tuple of (export_name, kernel_name), - where `export_name` is the name used to called from Python, - and `kernel_name` is the name of the kernel class in C++/CUDA source. - - :param args: Unique marker of the JIT module. Must be distinct for different kernels. - :type args: str - :param cpp_files: A list of C++ source files. - :type cpp_files: List[str] | None - :param cuda_files: A list of CUDA source files. - :type cuda_files: List[str] | None - :param cpp_wrappers: A list of C++ wrappers, defining the export name and kernel name. - :type cpp_wrappers: List[Tuple[str, str]] | None - :param cuda_wrappers: A list of CUDA wrappers, defining the export name and kernel name. - :type cuda_wrappers: List[Tuple[str, str]] | None - :param extra_cflags: Extra C++ compiler flags. - :type extra_cflags: List[str] | None - :param extra_cuda_cflags: Extra CUDA compiler flags. - :type extra_cuda_cflags: List[str] | None - :param extra_ldflags: Extra linker flags. - :type extra_ldflags: List[str] | None - :param extra_include_paths: Extra include paths. - :type extra_include_paths: List[str] | None - :param extra_dependencies: Extra dependencies for the JIT module, e.g., cutlass. - :type extra_dependencies: List[str] | None - :param build_directory: The build directory for JIT compilation. - :type build_directory: str | None - :param header_only: Whether the module is header-only. - If true, apply the wrappers to export given class/functions. - Otherwise, we must export from C++/CUDA side. - :return: A just-in-time(JIT) compiled module. - :rtype: Module - """ - - from tvm_ffi.cpp import load, load_inline - - cpp_files = cpp_files or [] - cuda_files = cuda_files or [] - extra_cflags = extra_cflags or [] - extra_cuda_cflags = extra_cuda_cflags or [] - extra_ldflags = extra_ldflags or [] - extra_include_paths = extra_include_paths or [] - - cpp_files = [str((KERNEL_PATH / "csrc" / f).resolve()) for f in cpp_files] - cuda_files = [str((KERNEL_PATH / "csrc" / f).resolve()) for f in cuda_files] - - for dep in set(extra_dependencies or []): - if dep not in _REGISTERED_DEPENDENCIES: - raise ValueError(f"Dependency {dep} is not registered.") - extra_include_paths += _REGISTERED_DEPENDENCIES[dep]() - - module_name = "sgl_kernel_jit_" + "_".join(str(arg) for arg in args) - if header_only: - cpp_wrappers = cpp_wrappers or [] - cuda_wrappers = cuda_wrappers or [] - cpp_sources = [f'#include "{path}"' for path in cpp_files] - cpp_sources += [_make_wrapper(tup) for tup in cpp_wrappers] - - # include cuda files - cuda_sources = [f'#include "{path}"' for path in cuda_files] - cuda_sources += [_make_wrapper(tup) for tup in cuda_wrappers] - with _jit_compile_context(): - return load_inline( - module_name, - cpp_sources=cpp_sources, - cuda_sources=cuda_sources, - extra_cflags=DEFAULT_CFLAGS + extra_cflags, - extra_cuda_cflags=_get_default_target_flags() + extra_cuda_cflags, - extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags, - extra_include_paths=DEFAULT_INCLUDE + extra_include_paths, - build_directory=build_directory, - ) - else: - assert cpp_wrappers is None and cuda_wrappers is None - with _jit_compile_context(): - return load( - module_name, - cpp_files=cpp_files, - cuda_files=cuda_files, - extra_cflags=DEFAULT_CFLAGS + extra_cflags, - extra_cuda_cflags=_get_default_target_flags() + extra_cuda_cflags, - extra_ldflags=DEFAULT_LDFLAGS + extra_ldflags, - extra_include_paths=DEFAULT_INCLUDE + extra_include_paths, - build_directory=build_directory, - ) - - -@dataclass -class ArchInfo: - major: int - minor: int - suffix: str - - @property - def target_name(self) -> str: - return f"{self.major}.{self.minor}{self.suffix}" - - @property - def jit_flag(self) -> str: - return f"-DSGL_CUDA_ARCH={self.major * 100 + self.minor * 10}" - - -@cache_once -def _init_jit_cuda_arch_once(): - global _CUDA_ARCH - try: - device = torch.cuda.current_device() - major, minor = torch.cuda.get_device_capability(device) - except Exception: - logger.warning("Cannot detect CUDA architecture.") - major, minor = 0, 0 # invalid value to trigger compile error if used - _CUDA_ARCH = ArchInfo(major, minor, "") - - -@contextmanager -def _jit_compile_context(): - if is_hip_runtime(): - yield # TODO: support ROCm `TVM_FFI_ROCM_ARCH_LIST` if needed - return - env_key = "TVM_FFI_CUDA_ARCH_LIST" - old_value = os.environ.get(env_key, None) - os.environ[env_key] = get_jit_cuda_arch().target_name - try: - yield - finally: - if old_value is None: - os.environ.pop(env_key, None) - else: - os.environ[env_key] = old_value - - -# NOTE: this might also be used in __main__.py for compile flags export -def _get_default_target_flags() -> List[str]: - if is_hip_runtime(): - flags = ["-DUSE_ROCM", "-std=c++20", "-O3"] - # Detect FP8 type based on GPU architecture - try: - device = torch.cuda.current_device() - gcn_arch = torch.cuda.get_device_properties(device).gcnArchName - if "gfx942" in gcn_arch: - flags.append("-DHIP_FP8_TYPE_FNUZ=1") - else: - flags.append("-DHIP_FP8_TYPE_E4M3=1") - except Exception: - flags.append("-DHIP_FP8_TYPE_E4M3=1") - return flags - else: - return [ - get_jit_cuda_arch().jit_flag, - "-std=c++20", - "-O3", - "--expt-relaxed-constexpr", - ] - - -@contextmanager -def override_jit_cuda_arch(major: int, minor: int, suffix: str = ""): - """A context manager to temporarily override CUDA architecture.""" - global _CUDA_ARCH - old_value = get_jit_cuda_arch() - _CUDA_ARCH = ArchInfo(major, minor, suffix) - try: - yield - finally: - _CUDA_ARCH = old_value - - -def get_jit_cuda_arch() -> ArchInfo: - """Get the current CUDA architecture info.""" - _init_jit_cuda_arch_once() - return _CUDA_ARCH - - -@cache_once -def is_arch_support_pdl() -> bool: - if is_hip_runtime() or is_musa_runtime(): - return False - return get_jit_cuda_arch().major >= 9 - - -def _find_package_root(package: str) -> Optional[pathlib.Path]: - spec = importlib.util.find_spec(package) - if spec is None or spec.origin is None: - return None - return pathlib.Path(spec.origin).resolve().parent - - -# NOTE: this might also be used in __main__.py for compile flags export -_REGISTERED_DEPENDENCIES: Dict[str, Callable[[], List[str]]] = {} - - -def register_dependency(name: str): - def decorator(f: Callable[[], List[str]]) -> Callable[[], List[str]]: - if name in _REGISTERED_DEPENDENCIES: - raise ValueError(f"Dependency {name} already registered") - _REGISTERED_DEPENDENCIES[name] = f - return f - - return decorator - - -@register_dependency("flashinfer") -def get_flashinfer_include_paths() -> List[str]: - include_paths: List[str] = [] - flashinfer_root = _find_package_root("flashinfer") - if flashinfer_root is None: - raise RuntimeError( - "Cannot find flashinfer package. Please install flashinfer to get" - "the required headers for JIT compilation." - ) - - flashinfer_data = flashinfer_root / "data" - candidates = [ - flashinfer_data / "include", - flashinfer_data / "csrc", - flashinfer_data / "cutlass" / "include", - flashinfer_data / "cutlass" / "tools" / "util" / "include", - flashinfer_data / "spdlog" / "include", - ] - - for path in candidates: - if not path.exists(): - raise RuntimeError( - f"Required header path {path} for flashinfer dependency not found." - " Please check your flashinfer installation." - ) - include_paths.append(str(path)) - return include_paths - - -@register_dependency("cutlass") -def get_cutlass_include_paths() -> List[str]: - include_paths: List[str] = [] - - flashinfer_root = _find_package_root("flashinfer") - if flashinfer_root is not None: - candidates = [ - flashinfer_root / "data" / "cutlass" / "include", - flashinfer_root / "data" / "cutlass" / "tools" / "util" / "include", - ] - for path in candidates: - if path.exists(): - include_paths.append(str(path)) - - deep_gemm_root = _find_package_root("deep_gemm") - if deep_gemm_root is not None: - candidate = deep_gemm_root / "include" - if candidate.exists(): - include_paths.append(str(candidate)) - - # De-duplicate while preserving order. - unique_paths = [] - seen = set() - for path in include_paths: - if path in seen: - continue - seen.add(path) - unique_paths.append(path) - - if not unique_paths: - raise RuntimeError( - "Cannot find CUTLASS headers required for JIT compilation. " - "Please install flashinfer or deep_gemm with CUTLASS headers." - ) - return unique_paths - - -__all__ = [ - "should_run_full_tests", - "get_ci_test_range", - "cache_once", - "is_hip_runtime", - "make_cpp_args", - "load_jit", - "override_jit_cuda_arch", - "get_jit_cuda_arch", - "is_arch_support_pdl", - "register_dependency", -] diff --git a/lightllm/third_party/sglang_jit/runtime_utils.py b/lightllm/third_party/sglang_jit/runtime_utils.py deleted file mode 100644 index d322498ca4..0000000000 --- a/lightllm/third_party/sglang_jit/runtime_utils.py +++ /dev/null @@ -1,5 +0,0 @@ -import torch - - -def is_hip() -> bool: - return torch.version.hip is not None From b39625c3de45ff39d8ac01779868ea53e31081d4 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 14 Jul 2026 02:40:12 +0000 Subject: [PATCH 077/214] optimize flashmla --- .../basemodel/attention/create_utils.py | 16 +- .../attention/nsa/dsv4_fp8_flashmla_sparse.py | 186 ++++++++++ .../attention/nsa/fp8_flashmla_sparse.py | 351 ++---------------- .../layer_infer/transformer_layer_infer.py | 1 - .../layer_infer/transformer_layer_infer.py | 26 +- lightllm/models/deepseek_v4/model.py | 29 +- lightllm/models/deepseek_v4/workspace.py | 15 + 7 files changed, 282 insertions(+), 342 deletions(-) create mode 100644 lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py diff --git a/lightllm/common/basemodel/attention/create_utils.py b/lightllm/common/basemodel/attention/create_utils.py index 594e81a9b4..95493d1de2 100644 --- a/lightllm/common/basemodel/attention/create_utils.py +++ b/lightllm/common/basemodel/attention/create_utils.py @@ -130,21 +130,25 @@ def get_mla_decode_att_backend_class(index=0, priority_list: list = ["flashinfer return _auto_select_backend(llm_dtype, kv_type_to_backend=mla_data_type_to_backend, priority_list=priority_list) -def get_nsa_prefill_att_backend_class(index=0, priority_list: list = ["flashmla_sparse"]) -> BaseAttBackend: +def get_nsa_prefill_att_backend_class( + index=0, priority_list: list = ["flashmla_sparse"], backend_map=nsa_data_type_to_backend +) -> BaseAttBackend: args = get_env_start_args() llm_dtype = args.llm_kv_type backend_str = args.llm_prefill_att_backend[index] if backend_str != "auto": - return nsa_data_type_to_backend[llm_dtype][backend_str] + return backend_map[llm_dtype][backend_str] else: - return _auto_select_backend(llm_dtype, kv_type_to_backend=nsa_data_type_to_backend, priority_list=priority_list) + return _auto_select_backend(llm_dtype, kv_type_to_backend=backend_map, priority_list=priority_list) -def get_nsa_decode_att_backend_class(index=0, priority_list: list = ["flashmla_sparse"]) -> BaseAttBackend: +def get_nsa_decode_att_backend_class( + index=0, priority_list: list = ["flashmla_sparse"], backend_map=nsa_data_type_to_backend +) -> BaseAttBackend: args = get_env_start_args() llm_dtype = args.llm_kv_type backend_str = args.llm_decode_att_backend[index] if backend_str != "auto": - return nsa_data_type_to_backend[llm_dtype][backend_str] + return backend_map[llm_dtype][backend_str] else: - return _auto_select_backend(llm_dtype, kv_type_to_backend=nsa_data_type_to_backend, priority_list=priority_list) + return _auto_select_backend(llm_dtype, kv_type_to_backend=backend_map, priority_list=priority_list) diff --git a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py new file mode 100644 index 0000000000..eece2648fb --- /dev/null +++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py @@ -0,0 +1,186 @@ +import dataclasses +from typing import TYPE_CHECKING + +import torch + +from ..base_att import AttControl, BaseAttBackend, BaseDecodeAttState, BasePrefillAttState + +if TYPE_CHECKING: + from lightllm.common.basemodel.infer_struct import InferStateInfo + + +# The current FlashMLA MODEL1 binary only instantiates these Q-head counts. +_SUPPORTED_Q_HEADS = (64, 128) + + +def get_dsv4_flashmla_padded_q_heads(q_head_num: int) -> int: + for supported_head_num in _SUPPORTED_Q_HEADS: + if q_head_num <= supported_head_num: + return supported_head_num + raise ValueError(f"FlashMLA does not support {q_head_num} local Q heads; supported counts: {_SUPPORTED_Q_HEADS}") + + +def _view_cache(buffer: torch.Tensor, page_size: int) -> torch.Tensor: + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_MLA_BYTES_PER_TOKEN + + byte_num = page_size * DSV4_MLA_BYTES_PER_TOKEN + return buffer[:, :byte_num].view(buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) + + +def _get_vllm_flashmla(): + from vllm.v1.attention.ops import flashmla + + return flashmla + + +class DeepseekV4FlashMlaFp8SparseAttBackend(BaseAttBackend): + def __init__(self, model): + super().__init__(model=model) + self.real_q_head_num = model.config["num_attention_heads"] // model.tp_world_size_ + self.padded_q_head_num = get_dsv4_flashmla_padded_q_heads(self.real_q_head_num) + self.compress_ratios = tuple(dict.fromkeys(model.config["compress_ratios"])) + + def _flashmla_att( + self, + q: torch.Tensor, + packed_kv: torch.Tensor, + mem_manager, + nsa_dict: dict, + sched_meta, + flash_mla, + flashmla_out: torch.Tensor = None, + ) -> torch.Tensor: + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_C128_PAGE_SIZE, + DSV4_C4_PAGE_SIZE, + DSV4_SWA_PAGE_SIZE, + ) + + ratio = nsa_dict["compress_ratio"] + extra_cache = None + if ratio == 4: + extra_page_size = DSV4_C4_PAGE_SIZE + elif ratio == 128: + extra_page_size = DSV4_C128_PAGE_SIZE + elif ratio != 0: + raise ValueError(f"unsupported DeepSeek-V4 compress ratio: {ratio}") + if ratio: + buffer = mem_manager.get_compressed_kv_buffer(nsa_dict["layer_index"]) + extra_cache = _view_cache(buffer, extra_page_size) + + kwargs = dict( + q=q.unsqueeze(1), + k_cache=_view_cache(packed_kv, DSV4_SWA_PAGE_SIZE), + block_table=None, + cache_seqlens=None, + head_dim_v=nsa_dict["head_dim_v"], + tile_scheduler_metadata=sched_meta, + num_splits=None, + softmax_scale=nsa_dict["softmax_scale"], + causal=False, + is_fp8_kvcache=True, + indices=nsa_dict["swa_indices"], + attn_sink=nsa_dict["attn_sink"], + topk_length=nsa_dict["swa_lengths"], + extra_k_cache=extra_cache, + extra_indices_in_kvcache=nsa_dict.get("extra_indices"), + extra_topk_length=nsa_dict.get("extra_lengths"), + ) + if flashmla_out is not None: + kwargs["out"] = flashmla_out + full_out, _ = flash_mla.flash_mla_with_kvcache(**kwargs) + return full_out[:, 0, : self.real_q_head_num, :] + + def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "_PrefillAttState": + return _PrefillAttState(backend=self, infer_state=infer_state) + + def create_att_decode_state(self, infer_state: "InferStateInfo") -> "_DecodeAttState": + return _DecodeAttState(backend=self, infer_state=infer_state) + + +@dataclasses.dataclass +class _PrefillAttState(BasePrefillAttState): + flashmla_sched_meta: dict = None + flash_mla: object = None + + def init_state(self): + self.flash_mla = _get_vllm_flashmla() + self.flashmla_sched_meta = {} + + def _get_sched_meta(self, compress_ratio: int): + if compress_ratio not in self.flashmla_sched_meta: + self.flashmla_sched_meta[compress_ratio] = self.flash_mla.get_mla_metadata()[0] + return self.flashmla_sched_meta[compress_ratio] + + def prefill_att( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + att_control: AttControl = AttControl(), + alloc_func=torch.empty, + *, + out: torch.Tensor = None, + ) -> torch.Tensor: + assert att_control.nsa_prefill, "nsa_prefill must be True for NSA prefill attention" + assert att_control.nsa_prefill_dict is not None, "nsa_prefill_dict is required" + nsa_dict = att_control.nsa_prefill_dict + if out is None: + out = alloc_func( + (q.shape[0], self.backend.real_q_head_num, nsa_dict["head_dim_v"]), + dtype=q.dtype, + device=q.device, + ) + full_out = self.infer_state.dsv4_workspace.flashmla_prefill_full_out[: q.shape[0]] + out.copy_( + self.backend._flashmla_att( + q, + k, + self.infer_state.mem_manager, + nsa_dict, + self._get_sched_meta(nsa_dict["compress_ratio"]), + self.flash_mla, + flashmla_out=full_out, + ) + ) + return out + + +@dataclasses.dataclass +class _DecodeAttState(BaseDecodeAttState): + flashmla_sched_meta: dict = None + flash_mla: object = None + + def init_state(self): + self.reset_sched_meta_for_capture() + + def reset_sched_meta_for_capture(self): + # FlashMLA lazily binds extra-cache geometry, so ratios cannot share one sched-meta object. + self.flash_mla = _get_vllm_flashmla() + self.flashmla_sched_meta = { + ratio: self.flash_mla.get_mla_metadata()[0] for ratio in self.backend.compress_ratios + } + + def decode_att( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + att_control: AttControl = AttControl(), + alloc_func=torch.empty, + ) -> torch.Tensor: + assert att_control.nsa_decode, "nsa_decode must be True for NSA decode attention" + assert att_control.nsa_decode_dict is not None, "nsa_decode_dict is required" + nsa_dict = att_control.nsa_decode_dict + real_out = self.backend._flashmla_att( + q, + k, + self.infer_state.mem_manager, + nsa_dict, + self.flashmla_sched_meta[nsa_dict["compress_ratio"]], + self.flash_mla, + ) + return real_out.contiguous() + + +DSV4_NSA_BACKENDS = {"fp8kv_dsa": {"flashmla_sparse": DeepseekV4FlashMlaFp8SparseAttBackend}} diff --git a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py index c0ba1d027b..539ade769e 100644 --- a/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/fp8_flashmla_sparse.py @@ -1,155 +1,22 @@ import dataclasses -import inspect import torch from typing import TYPE_CHECKING, Tuple from ..base_att import AttControl, BaseAttBackend, BaseDecodeAttState, BasePrefillAttState from lightllm.utils.dist_utils import get_current_device_id -from lightllm.utils.log_utils import init_logger if TYPE_CHECKING: from lightllm.common.basemodel.infer_struct import InferStateInfo -logger = init_logger(__name__) - -# this flash_mla extra-cache fork only instantiates h_q in {64, 128}; pad TP-split q heads up -# to the nearest supported count (zero heads are discarded from the output slice). -FLASHMLA_SUPPORTED_HEADS = (64, 128) - - -def _target_q_heads(h_q: int) -> int: - target = next((h for h in FLASHMLA_SUPPORTED_HEADS if h >= h_q), None) - assert target is not None, f"num q heads {h_q} exceeds flash_mla support {FLASHMLA_SUPPORTED_HEADS}" - return target - - -def _pad_q_heads( - q_4d: torch.Tensor, - attn_sink: torch.Tensor, - q_out: torch.Tensor = None, - sink_out: torch.Tensor = None, -): - h_q = q_4d.shape[2] - if h_q in FLASHMLA_SUPPORTED_HEADS: - return q_4d, attn_sink, h_q - target = _target_q_heads(h_q) - if q_out is not None: - q_out[:, :, :h_q, :].copy_(q_4d) - q_out[:, :, h_q:target, :].zero_() - sink_out[:h_q].copy_(attn_sink) - sink_out[h_q:target].zero_() - return q_out, sink_out[:target], h_q - q_pad = torch.nn.functional.pad(q_4d, (0, 0, 0, target - h_q)) - sink_pad = torch.nn.functional.pad(attn_sink, (0, target - h_q)) - return q_pad, sink_pad, h_q - - -def _view_dsv4_flashmla_cache(layer_buffer: torch.Tensor, page_size: int) -> torch.Tensor: - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_MLA_BYTES_PER_TOKEN - - usable = page_size * DSV4_MLA_BYTES_PER_TOKEN - return layer_buffer[:, :usable].view(layer_buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) - - -@dataclasses.dataclass -class _Dsv4Metadata: - swa_indices: torch.Tensor - swa_lengths: torch.Tensor - extra_cache: torch.Tensor = None - extra_indices: torch.Tensor = None - extra_lengths: torch.Tensor = None - - -def _metadata_from_dict(infer_state, nsa_dict: dict) -> "_Dsv4Metadata": - """Bundle the model-built FINAL index tensors (carried in nsa_dict by DeepseekV4IndexInfer) with - the layer-keyed fp8 extra-cache byte view. The cache view is data-independent (a fixed per-layer - buffer slice), so it is built here -- a genuine flash_mla ABI concern -- rather than on the model - side; only the index/length tensors cross the att_control boundary.""" - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C128_PAGE_SIZE, DSV4_C4_PAGE_SIZE - - ratio = nsa_dict["compress_ratio"] - extra_cache = None - if ratio: - page = DSV4_C4_PAGE_SIZE if ratio == 4 else DSV4_C128_PAGE_SIZE - extra_buffer = infer_state.mem_manager.get_compressed_kv_buffer(nsa_dict["layer_index"]) - extra_cache = _view_dsv4_flashmla_cache(extra_buffer, page) - return _Dsv4Metadata( - swa_indices=nsa_dict["swa_indices"], - swa_lengths=nsa_dict["swa_lengths"], - extra_cache=extra_cache, - extra_indices=nsa_dict.get("extra_indices"), - extra_lengths=nsa_dict.get("extra_lengths"), - ) - class NsaFlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) - self.use_dsv4_flashmla_kvcache = model.config.get("model_type") == "deepseek_v4" - if self.use_dsv4_flashmla_kvcache: - self.ragged_mem_buffers = None - logger.info("DSV4 FlashMLA kvcache path skips generic NSA ragged decode buffers") - else: - device = get_current_device_id() - self.ragged_mem_buffers = [ - torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) - for _ in range(2) - ] - self.prefill_flash_mla, self.prefill_flash_mla_supports_out = self._load_prefill_flash_mla() - self.prefill_q_workspace = None - self.prefill_out_workspace = None - self.prefill_real_out_workspace = None - self.prefill_sink_workspace = None - self.prefill_workspace_shape = None - self.prefill_real_out_shape = None - self.prefill_workspace_token_capacity = int(model.batch_max_tokens or 0) - if self.prefill_flash_mla_supports_out: - logger.info("DSV4 FlashMLA prefill uses vLLM out= workspace path") - else: - logger.warning("DSV4 FlashMLA prefill out= path unavailable; falling back to allocating FlashMLA output") - - def _load_prefill_flash_mla(self): - try: - from vllm.v1.attention.ops import flashmla as flash_mla - - sig = inspect.signature(flash_mla.flash_mla_with_kvcache) - if "out" in sig.parameters: - return flash_mla, True - except Exception: - pass - - import flash_mla - - return flash_mla, "out" in inspect.signature(flash_mla.flash_mla_with_kvcache).parameters - - def _ensure_prefill_workspace(self, token_num: int, head_num: int, target_heads: int, head_dim: int, dtype, device): - capacity = max(token_num, self.prefill_workspace_token_capacity) - workspace_shape = (capacity, 1, target_heads, head_dim) - if ( - self.prefill_workspace_shape != (target_heads, head_dim, dtype, device) - or self.prefill_q_workspace is None - or self.prefill_q_workspace.shape[0] < capacity - ): - self.prefill_q_workspace = torch.empty(workspace_shape, dtype=dtype, device=device) - self.prefill_out_workspace = torch.empty(workspace_shape, dtype=dtype, device=device) - self.prefill_sink_workspace = torch.empty((target_heads,), dtype=torch.float32, device=device) - self.prefill_workspace_shape = (target_heads, head_dim, dtype, device) - - real_out_shape = (capacity, head_num, head_dim) - if ( - self.prefill_real_out_shape != (head_num, head_dim, dtype, device) - or self.prefill_real_out_workspace is None - or self.prefill_real_out_workspace.shape[0] < capacity - ): - self.prefill_real_out_workspace = torch.empty(real_out_shape, dtype=dtype, device=device) - self.prefill_real_out_shape = (head_num, head_dim, dtype, device) - - return ( - self.prefill_q_workspace[:token_num], - self.prefill_out_workspace[:token_num], - self.prefill_real_out_workspace[:token_num], - self.prefill_sink_workspace, - ) + device = get_current_device_id() + self.ragged_mem_buffers = [ + torch.empty(model.graph_max_batch_size * model.max_seq_length, dtype=torch.int32, device=device) + for _ in range(2) + ] def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "NsaFlashMlaFp8SparsePrefillAttState": return NsaFlashMlaFp8SparsePrefillAttState(backend=self, infer_state=infer_state) @@ -164,20 +31,9 @@ class NsaFlashMlaFp8SparsePrefillAttState(BasePrefillAttState): ke: torch.Tensor = None lengths: torch.Tensor = None ragged_mem_index: torch.Tensor = None - flashmla_sched_meta: object = None def init_state(self): self.backend: NsaFlashMlaFp8SparseAttBackend = self.backend - self.flashmla_sched_meta = {} - return - - def ensure_nsa_ks_ke(self): - """Build the ragged ks/ke/lengths (+ ragged_mem_index) the DeepSeek-3.2 indexer consumes. The - indexer calls this explicitly before reading them; DeepSeek-V4 uses its own indexer and never - calls it, so V4 prefill skips the alloc + gen_nsa_ks_ke kernel. Idempotent + layer-independent: - the first call in a forward computes, the other layers reuse.""" - if self.ks is not None: - return self.ragged_mem_index = torch.empty( self.infer_state.total_token_num, dtype=torch.int32, @@ -196,13 +52,6 @@ def ensure_nsa_ks_ke(self): ) return - def _get_flashmla_sched_meta(self, compress_ratio: int): - sched_meta = self.flashmla_sched_meta.get(compress_ratio) - if sched_meta is None: - sched_meta = self.backend.prefill_flash_mla.get_mla_metadata()[0] - self.flashmla_sched_meta[compress_ratio] = sched_meta - return sched_meta - def prefill_att( self, q: torch.Tensor, @@ -210,17 +59,9 @@ def prefill_att( v: torch.Tensor, att_control: AttControl = AttControl(), alloc_func=torch.empty, - out: torch.Tensor = None, ) -> torch.Tensor: assert att_control.nsa_prefill, "nsa_prefill must be True for NSA prefill attention" assert att_control.nsa_prefill_dict is not None, "nsa_prefill_dict is required" - if att_control.nsa_prefill_dict.get("flashmla_kvcache"): - return self._flashmla_kvcache_prefill_att( - q=q, - packed_kv=k, - nsa_dict=att_control.nsa_prefill_dict, - out=out, - ) return self._nsa_prefill_att(q=q, packed_kv=k, att_control=att_control) def _nsa_prefill_att( @@ -237,8 +78,6 @@ def _nsa_prefill_att( kv_lora_rank = nsa_dict["kv_lora_rank"] topk_mem_indices = nsa_dict["topk_mem_indices"] prefill_cache_kv = nsa_dict["prefill_cache_kv"] - attn_sink = nsa_dict.get("attn_sink") - topk_length = nsa_dict.get("topk_length") if self.infer_state.prefix_total_token_num > 0: # 当前推理生成的token kv部分从 prefill_cache_kv 中获取,历史 @@ -262,69 +101,9 @@ def _nsa_prefill_att( indices=topk_indices, sm_scale=softmax_scale, d_v=kv_lora_rank, - attn_sink=attn_sink, - topk_length=topk_length, ) return mla_out - def _flashmla_kvcache_prefill_att( - self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict, out: torch.Tensor = None - ) -> torch.Tensor: - attn_sink = nsa_dict["attn_sink"] - metadata = _metadata_from_dict(self.infer_state, nsa_dict) - return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict, out=out) - - def _flashmla_kvcache_att( - self, - q: torch.Tensor, - packed_kv: torch.Tensor, - metadata: _Dsv4Metadata, - attn_sink: torch.Tensor, - nsa_dict: dict, - out: torch.Tensor = None, - ) -> torch.Tensor: - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE - - q_4d = q.unsqueeze(1).contiguous() - num_real_heads = q_4d.shape[2] - target_heads = _target_q_heads(num_real_heads) - q_workspace, full_out_workspace, real_out_workspace, sink_workspace = self.backend._ensure_prefill_workspace( - q_4d.shape[0], - num_real_heads, - target_heads, - q_4d.shape[-1], - q_4d.dtype, - q_4d.device, - ) - q_for_flash, sink_for_flash, num_real_heads = _pad_q_heads(q_4d, attn_sink, q_workspace, sink_workspace) - k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) - sched_meta = self._get_flashmla_sched_meta(nsa_dict["compress_ratio"]) - flash_mla = self.backend.prefill_flash_mla - kwargs = dict( - q=q_for_flash, - k_cache=k_cache, - block_table=None, - cache_seqlens=None, - head_dim_v=nsa_dict["head_dim_v"], - tile_scheduler_metadata=sched_meta, - num_splits=None, - softmax_scale=nsa_dict["softmax_scale"], - causal=False, - is_fp8_kvcache=True, - indices=metadata.swa_indices, - attn_sink=sink_for_flash, - topk_length=metadata.swa_lengths, - extra_k_cache=metadata.extra_cache, - extra_indices_in_kvcache=metadata.extra_indices, - extra_topk_length=metadata.extra_lengths, - ) - if self.backend.prefill_flash_mla_supports_out: - kwargs["out"] = full_out_workspace - full_out, _ = flash_mla.flash_mla_with_kvcache(**kwargs) - real_out = out if out is not None else real_out_workspace - real_out.copy_(full_out[:, 0, :num_real_heads, :]) - return real_out - @dataclasses.dataclass class NsaFlashMlaFp8SparseDecodeAttState(BaseDecodeAttState): @@ -336,52 +115,35 @@ class NsaFlashMlaFp8SparseDecodeAttState(BaseDecodeAttState): def init_state(self): self.backend: NsaFlashMlaFp8SparseAttBackend = self.backend - if not self.backend.use_dsv4_flashmla_kvcache: - model = self.backend.model - use_cuda_graph = ( - self.infer_state.batch_size <= model.graph_max_batch_size - and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch - ) - - if use_cuda_graph: - self.ragged_mem_index = self.backend.ragged_mem_buffers[self.infer_state.microbatch_index] - else: - self.ragged_mem_index = torch.empty( - self.infer_state.total_token_num, - dtype=torch.int32, - device=get_current_device_id(), - ) - - from lightllm.common.basemodel.triton_kernel.gen_nsa_ks_ke import gen_nsa_ks_ke + model = self.backend.model + use_cuda_graph = ( + self.infer_state.batch_size <= model.graph_max_batch_size + and self.infer_state.max_kv_seq_len <= model.graph_max_len_in_batch + ) - self.ks, self.ke, self.lengths = gen_nsa_ks_ke( - b_seq_len=self.infer_state.b_seq_len, - 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.b_seq_len.shape[0], - ragged_mem_index=self.ragged_mem_index, - hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, + if use_cuda_graph: + self.ragged_mem_index = self.backend.ragged_mem_buffers[self.infer_state.microbatch_index] + else: + self.ragged_mem_index = torch.empty( + self.infer_state.total_token_num, + dtype=torch.int32, + device=get_current_device_id(), ) - import flash_mla - # one sched_meta per layer type: the lazy config locks extra-cache geometry (page size, - # presence) on first invocation, so swa-only/c4/c128 layers must not share one object. - self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} - return - - def ensure_nsa_ks_ke(self): - # decode builds ks/ke eagerly in init_state (outside the cuda graph, for capture safety), so - # they are already available -- this satisfies the shared DeepSeek-3.2 indexer ensure contract. - return + from lightllm.common.basemodel.triton_kernel.gen_nsa_ks_ke import gen_nsa_ks_ke - def reset_sched_meta_for_capture(self): - # cuda-graph capture hook: the warmup pass already locked/stored sched meta on this - # (shared) state object; reset so the capture pass re-plans INSIDE the graph and every - # replay re-plans from the live tensors instead of binding warmup leftovers. + self.ks, self.ke, self.lengths = gen_nsa_ks_ke( + b_seq_len=self.infer_state.b_seq_len, + 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.b_seq_len.shape[0], + ragged_mem_index=self.ragged_mem_index, + hold_req_idx=self.infer_state.req_manager.HOLD_REQUEST_ID, + ) import flash_mla - self.flashmla_sched_meta = {ratio: flash_mla.get_mla_metadata()[0] for ratio in (0, 4, 128)} + self.flashmla_sched_meta, _ = flash_mla.get_mla_metadata() return def decode_att( @@ -394,12 +156,6 @@ def decode_att( ) -> torch.Tensor: assert att_control.nsa_decode, "nsa_decode must be True for NSA decode attention" assert att_control.nsa_decode_dict is not None, "nsa_decode_dict is required" - if att_control.nsa_decode_dict.get("flashmla_kvcache"): - return self._flashmla_kvcache_decode_att( - q=q, - packed_kv=k, - nsa_dict=att_control.nsa_decode_dict, - ) return self._nsa_decode_att(q=q, packed_kv=k, att_control=att_control) def _nsa_decode_att( @@ -414,11 +170,6 @@ def _nsa_decode_att( topk_mem_indices = nsa_dict["topk_mem_indices"] softmax_scale = nsa_dict["softmax_scale"] kv_lora_rank = nsa_dict["kv_lora_rank"] - attn_sink = nsa_dict.get("attn_sink") - topk_length = nsa_dict.get("topk_length") - extra_k_cache = nsa_dict.get("extra_k_cache") - extra_indices = nsa_dict.get("extra_indices_in_kvcache") - extra_topk_length = nsa_dict.get("extra_topk_length") if topk_mem_indices.ndim == 2: topk_mem_indices = topk_mem_indices.unsqueeze(1) @@ -438,54 +189,10 @@ def _nsa_decode_att( block_table=None, cache_seqlens=None, head_dim_v=kv_lora_rank, - tile_scheduler_metadata=self.flashmla_sched_meta[0], + tile_scheduler_metadata=self.flashmla_sched_meta, softmax_scale=softmax_scale, causal=False, is_fp8_kvcache=True, indices=topk_mem_indices, - attn_sink=attn_sink, - topk_length=topk_length, - extra_k_cache=extra_k_cache, - extra_indices_in_kvcache=extra_indices, - extra_topk_length=extra_topk_length, ) return o_tensor[:, 0, :, :] # [b, 1, h, d] -> [b, h, d] - - def _flashmla_kvcache_decode_att(self, q: torch.Tensor, packed_kv: torch.Tensor, nsa_dict: dict) -> torch.Tensor: - attn_sink = nsa_dict["attn_sink"] - metadata = _metadata_from_dict(self.infer_state, nsa_dict) - return self._flashmla_kvcache_att(q, packed_kv, metadata, attn_sink, nsa_dict) - - def _flashmla_kvcache_att( - self, - q: torch.Tensor, - packed_kv: torch.Tensor, - metadata: _Dsv4Metadata, - attn_sink: torch.Tensor, - nsa_dict: dict, - ) -> torch.Tensor: - import flash_mla - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE - - q_4d = q.unsqueeze(1).contiguous() - q_4d, attn_sink, num_real_heads = _pad_q_heads(q_4d, attn_sink) - k_cache = _view_dsv4_flashmla_cache(packed_kv, DSV4_SWA_PAGE_SIZE) - out, _ = flash_mla.flash_mla_with_kvcache( - q=q_4d, - k_cache=k_cache, - block_table=None, - cache_seqlens=None, - head_dim_v=nsa_dict["head_dim_v"], - tile_scheduler_metadata=self.flashmla_sched_meta[nsa_dict["compress_ratio"]], - num_splits=None, - softmax_scale=nsa_dict["softmax_scale"], - causal=False, - is_fp8_kvcache=True, - indices=metadata.swa_indices, - attn_sink=attn_sink, - topk_length=metadata.swa_lengths, - extra_k_cache=metadata.extra_cache, - extra_indices_in_kvcache=metadata.extra_indices, - extra_topk_length=metadata.extra_lengths, - ) - return out[:, 0, :num_real_heads].contiguous() diff --git a/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py index 8c7506b1e1..d6eaebe2fd 100644 --- a/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek3_2/layer_infer/transformer_layer_infer.py @@ -206,7 +206,6 @@ def _get_indices( weights = layer_weight.weights_proj_.mm(hidden_states) * self.index_n_heads_scale weights = weights.unsqueeze(-1) * q_scale - att_state.ensure_nsa_ks_ke() ks = att_state.ks ke = att_state.ke lengths = att_state.lengths diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 3e0315bcb5..b7e125e547 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -51,6 +51,7 @@ def __init__(self, layer_num, network_config): self.swiglu_limit = float(network_config["swiglu_limit"]) self.softmax_scale = (self.qk_nope_head_dim + self.qk_rope_head_dim) ** (-0.5) self.tp_q_head_num_ = self.num_heads // self.tp_world_size_ + self.flashmla_q_head_num_ = self.tp_q_head_num_ self.tp_groups = self.o_groups // self.tp_world_size_ self.enable_ep_moe = get_env_start_args().enable_ep_moe self.compressor = CompressorInfer( @@ -162,8 +163,15 @@ def _get_qkv( q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, # fused in one DSV4 CUDA kernel (fp32 norm/rotation, bf16 in between -- same as eager). - q = self.alloc_tensor(q_in.shape, dtype=q_in.dtype, device=q_in.device) - fused_q_norm_rope(q_in, q, self.eps_, self.freqs_cis, infer_state.position_ids) + # The selected FlashMLA MODEL1 binary only instantiates H=64/128. Produce that ABI layout + # directly. Prefill uses one max-token workspace so changing T only changes the prefix view; + # decode keeps its graph-owned tensor. The workspace's padded tail is zeroed once at init. + if infer_state.is_prefill: + q = infer_state.dsv4_workspace.flashmla_prefill_q[:T] + else: + q = self.alloc_tensor((T, self.flashmla_q_head_num_, self.head_dim_), dtype=q_in.dtype, device=q_in.device) + q[:, self.tp_q_head_num_ :, :].zero_() + fused_q_norm_rope(q_in, q[:, : self.tp_q_head_num_, :], self.eps_, self.freqs_cis, infer_state.position_ids) # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 DSV4 CUDA kernel 完成, # 替代 eager norm/rope/cat + _post_cache_kv。 # bf16 kv 中间量没有其他消费者: flashmla 路径注意力读 cache,压缩器/indexer 取 x。 @@ -225,10 +233,7 @@ def _context_attention_wrapper_run( _o = tensor_to_no_ref_tensor(o) def att_func(new_infer_state: DeepseekV4InferStateInfo): - tmp_o = self._context_attention_kernel(_q, _q_lora, _x, new_infer_state, layer_weight, out=_o) - assert tmp_o.shape == _o.shape - if tmp_o.data_ptr() != _o.data_ptr(): - _o.copy_(tmp_o) + self._context_attention_kernel(_q, _q_lora, _x, new_infer_state, layer_weight, out=_o) return infer_state.prefill_cuda_graph_add_cpu_runnning_func(func=att_func, after_graph=pre_capture_graph) @@ -283,7 +288,6 @@ def _context_attention_kernel( att_control = AttControl( nsa_prefill=True, nsa_prefill_dict={ - "flashmla_kvcache": True, "layer_index": self.layer_num_, "compress_ratio": self.compress_ratio, "head_dim_v": self.v_head_dim, @@ -292,18 +296,19 @@ def _context_attention_kernel( **meta, }, ) - out = infer_state.prefill_att_state.prefill_att( + attn_out = infer_state.prefill_att_state.prefill_att( q=q, k=infer_state.mem_manager.get_att_input_params(layer_index=self.layer_num_), v=None, att_control=att_control, + alloc_func=self.alloc_tensor, out=out, ) pad_q_len = getattr(infer_state, "_dsv4_prefill_pad_q_len", 0) if pad_q_len: # pad 行读 HOLD 槽位(参见 infer_struct._dsv4_prefill_pad_q_len),清零以保持确定性 - out[-pad_q_len:] = 0 - return out + attn_out[-pad_q_len:] = 0 + return attn_out # ------------------------------------------------------------------ attention (decode) def token_attention_forward( @@ -323,7 +328,6 @@ def _token_attention_kernel( att_control = AttControl( nsa_decode=True, nsa_decode_dict={ - "flashmla_kvcache": True, "layer_index": self.layer_num_, "compress_ratio": self.compress_ratio, "head_dim_v": self.v_head_dim, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 567b14ea49..fd13bc43a0 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -25,6 +25,7 @@ DeepseekV4TransformerLayerInfer, ) from lightllm.common.basemodel.attention import get_nsa_prefill_att_backend_class, get_nsa_decode_att_backend_class +from lightllm.common.basemodel.attention.nsa.dsv4_fp8_flashmla_sparse import DSV4_NSA_BACKENDS from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo from lightllm.models.deepseek_v4.workspace import DeepseekV4Workspace from lightllm.models.llama.yarn_rotary_utils import ( @@ -109,8 +110,32 @@ def _init_att_backend(self): # TODO: 支持其他 kv type if args.llm_kv_type != "fp8kv_dsa": raise RuntimeError("DeepSeek-V4 requires llm_kv_type=fp8kv_dsa for packed FlashMLA sparse attention") - self.prefill_att_backend = get_nsa_prefill_att_backend_class(index=0)(model=self) - self.decode_att_backend = get_nsa_decode_att_backend_class(index=0)(model=self) + self.prefill_att_backend = get_nsa_prefill_att_backend_class(index=0, backend_map=DSV4_NSA_BACKENDS)(model=self) + self.decode_att_backend = get_nsa_decode_att_backend_class(index=0, backend_map=DSV4_NSA_BACKENDS)(model=self) + + real_q_head_num = self.prefill_att_backend.real_q_head_num + padded_q_head_num = self.prefill_att_backend.padded_q_head_num + self.dsv4_workspace.init_flashmla_prefill_q( + real_q_head_num=real_q_head_num, + padded_q_head_num=padded_q_head_num, + head_dim=self.config["head_dim"], + dtype=self.data_type, + ) + self.dsv4_workspace.init_flashmla_prefill_full_out( + q_head_num=padded_q_head_num, + head_dim_v=self.config["head_dim"], + dtype=self.data_type, + ) + for layer_infer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): + layer_infer.flashmla_q_head_num_ = padded_q_head_num + if padded_q_head_num == real_q_head_num: + continue + attn_sink = layer_weight.attn_sink_.weight + assert attn_sink.shape == (real_q_head_num,) + padded_attn_sink = torch.zeros((padded_q_head_num,), dtype=attn_sink.dtype, device=attn_sink.device) + padded_attn_sink[: attn_sink.shape[0]].copy_(attn_sink) + padded_attn_sink.load_ok = attn_sink.load_ok + layer_weight.attn_sink_.weight = padded_attn_sink return def _init_custom(self): diff --git a/lightllm/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py index 562a53bbf3..2b260af69a 100644 --- a/lightllm/models/deepseek_v4/workspace.py +++ b/lightllm/models/deepseek_v4/workspace.py @@ -18,6 +18,21 @@ def __init__(self, model): self.c4_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") self.c128_indices = self._alloc(self.c128_cap) self.c128_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") + self.flashmla_prefill_q = None + self.flashmla_prefill_full_out = None + + def init_flashmla_prefill_q(self, real_q_head_num: int, padded_q_head_num: int, head_dim: int, dtype: torch.dtype): + if self.flashmla_prefill_q is None: + self.flashmla_prefill_q = torch.empty( + (self.token_capacity, padded_q_head_num, head_dim), dtype=dtype, device="cuda" + ) + self.flashmla_prefill_q[:, real_q_head_num:, :].zero_() + + def init_flashmla_prefill_full_out(self, q_head_num: int, head_dim_v: int, dtype: torch.dtype): + if self.flashmla_prefill_full_out is None: + self.flashmla_prefill_full_out = torch.empty( + (self.token_capacity, 1, q_head_num, head_dim_v), dtype=dtype, device="cuda" + ) @staticmethod def compress_cap(max_kv_seq_len: int, ratio: int) -> int: From 51194ee7d19ddad8aa7ca35f792f0f39808343e3 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 14 Jul 2026 02:41:03 +0000 Subject: [PATCH 078/214] fix 9.11 < 9.9 --- lightllm/server/function_call_parser.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index c08cb83fa3..a66e383a46 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1572,11 +1572,12 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami if partial_len: return StreamingParseResult() + normal_text = current_text self._buffer = "" for e_token in [self.eot_token, self.invoke_end_token]: - if e_token in new_text: - new_text = new_text.replace(e_token, "") - return StreamingParseResult(normal_text=new_text) + if e_token in normal_text: + normal_text = normal_text.replace(e_token, "") + return StreamingParseResult(normal_text=normal_text) # Mark that we're inside a function_calls block if self.has_tool_call(current_text): From 5f17bd4c34b441ea22ba4ff8a4d3eacc6baf1abc Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 14 Jul 2026 09:37:29 +0000 Subject: [PATCH 079/214] refact --- .../deepseek4_mem_manager.py | 37 ------------- lightllm/common/req_manager.py | 52 ++++--------------- lightllm/models/deepseek_v4/model.py | 11 +--- 3 files changed, 11 insertions(+), 89 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 883bb59937..d07406529b 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -150,7 +150,6 @@ def __init__( compress_rates: List[int], indexer_head_dim: int = 128, max_request_num: Optional[int] = None, - sliding_window: Optional[int] = None, swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, @@ -168,7 +167,6 @@ def __init__( self.n_c128 = sum(1 for r in self.compress_rates if r == 128) self.indexer_head_dim = indexer_head_dim self.max_request_num = max_request_num - self.sliding_window = sliding_window self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) @@ -252,7 +250,6 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.c128_size = _ceil_div(size, 128) self.c4_pool: Optional[PackedPagePool] = None self.c4_indexer_pool: Optional[PackedPagePool] = None - self.c4_allocator: Optional[KvCacheAllocator] = None self.c4_page_allocator: Optional[KvCacheAllocator] = None self.c4_page_live_count: Optional[torch.Tensor] = None self.c128_pool: Optional[PackedPagePool] = None @@ -577,15 +574,9 @@ def free_all(self): self.full_to_c128_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c128_pool.HOLD_TOKEN_MEMINDEX return - def alloc_c4(self, need_size) -> torch.Tensor: - raise AssertionError("DeepSeek-V4 c4 uses page-safe allocation; call alloc_c4_pages instead") - def alloc_c128(self, need_size) -> torch.Tensor: return self.c128_allocator.alloc(need_size) - def free_c4(self, free_index) -> None: - raise AssertionError("DeepSeek-V4 c4 uses page live-count release; call evict_c4 instead") - def free_c128(self, free_index) -> None: self.c128_allocator.free(free_index) @@ -735,19 +726,6 @@ def pack_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer self.c4_indexer_pool.page_size, ) - def gather_indexer_k(self, layer_index: int, slots: torch.Tensor) -> torch.Tensor: - """反量化 gather c4 indexer-K: slots [N](c4 槽位,HOLD 合法) -> [N, indexer_head_dim] bf16。 - indexer top-k 打分用(纯张量操作,cuda-graph 安全)。""" - assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" - pool = self.c4_indexer_pool - flat = pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]).view(-1) - data_offsets, scale_offsets = pool._loc_offsets(slots.reshape(-1)) - data_range = torch.arange(pool.data_bytes_per_token, device=flat.device) - scale_range = torch.arange(pool.scale_bytes_per_token, device=flat.device) - k_fp8 = flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)].view(torch.float8_e4m3fn) - scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)].contiguous().view(torch.float32) - return (k_fp8.float() * scale).to(torch.bfloat16) - # ------------------------------------------------------------------ fenced inherited APIs # kv_buffer 是 page 索引的 uint8 slab,基类按 token 索引读写的接口会静默写坏数据,显式拦截。 def get_index_kv_buffer(self, index): @@ -756,9 +734,6 @@ def get_index_kv_buffer(self, index): def load_index_kv_buffer(self, index, load_tensor_dict): raise NotImplementedError("DeepSeek-V4 packed page-slab cache does not support token-indexed kv_buffer io") - def alloc_kv_move_buffer(self, max_req_total_len): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") @@ -767,15 +742,3 @@ def write_mem_to_page_kv_move_buffer(self, *args, **kwargs): def read_page_kv_move_buffer_to_mem(self, *args, **kwargs): raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - - def send_to_decode_node(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - - def receive_from_prefill_node(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - - def send_to_decode_node_p2p(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") - - def receive_from_prefill_node_p2p(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index fe49e9a095..8eb44b60d7 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -368,20 +368,8 @@ def copy_small_page_buffer_to_linear_att_state( class DeepseekV4ReqManager(ReqManager): """DeepSeek-V4 的请求级管理。 - 在基类 ReqManager 之上补 V4 专有的 per-request 结构。该对象在 mem manager profile 前创建, - 所以初始化只依赖 config 派生出的 compress_rates/head_dim/indexer_head_dim/sliding_window; - 真实 mem_manager 会在 `_init_mem_manager()` 后通过 `bind_mem_manager()` 接入。 - - * 压缩槽位不在本类: ``full_to_c4/c128_indexs``(mem manager)以组末 token 的 full 槽位为键。 - 本类只负责 prep 阶段的分配与 scatter(``prepare_prefill`` / - ``prepare_decode_compress_slots``)——必须先于 attention metadata 构建/图捕获; - 条目内容由 layer-infer 的 compressor 前向写入。 - * compressor 在途状态不在本类: c4/c128 都在 mem manager 的 swa 页派生池, - 随页生灭,命中零拷贝续算。 - * SWA 槽位分配/出窗回收(``prepare_prefill_swa`` / ``prepare_decode_swa``): 每步 prep 阶段 - 为新 token 调 mem_manager.alloc_swa,并按 per-req 水位线(``_swa_evict_marks``)惰性回收 - 已出窗位置的 swa 槽。水位线首次置为该请求首个 chunk 的 ready_cache_len(radix 共享前缀 - 的边界),因此共享前缀的 swa 槽永远不会被本请求回收(归 radix 经 mem_manager.free 级联释放)。 + 负责 req/seq/MTP 布局、SWA 回收水位线和派生槽位准备;具体池结构、映射和分配器 + 由 ``DeepseekV4MemoryManager`` 持有。对象先于 mem manager 创建,模型初始化后再接入。 """ def __init__( @@ -389,9 +377,6 @@ def __init__( max_request_num, max_sequence_length, mem_manager: Optional[DeepseekV4MemoryManager] = None, - compress_rates: Optional[List[int]] = None, - head_dim: Optional[int] = None, - indexer_head_dim: Optional[int] = None, sliding_window: Optional[int] = None, ): super().__init__(max_request_num, max_sequence_length, mem_manager) @@ -400,23 +385,6 @@ def __init__( # 出窗回收水位线: -1 表示该 req 尚未见过 prefill chunk(首个 chunk 的 ready_cache_len # 即共享前缀边界,作为永不下探的回收下界)。 self._swa_evict_marks = [-1 for _ in range(max_request_num + 1)] - self.compress_rates = list(compress_rates) - self.n_c4 = sum(1 for r in self.compress_rates if r == 4) - self.n_c128 = sum(1 for r in self.compress_rates if r == 128) - self.head_dim = head_dim - self.indexer_head_dim = indexer_head_dim - self.layer_to_c4_idx = {} - self.layer_to_c128_idx = {} - self.mem_manager = mem_manager - c4 = c128 = 0 - for lid, r in enumerate(self.compress_rates): - if r == 4: - self.layer_to_c4_idx[lid] = c4 - c4 += 1 - elif r == 128: - self.layer_to_c128_idx[lid] = c128 - c128 += 1 - return # ------------------------------------------------------------------ swa slot prep (per step) @@ -741,7 +709,7 @@ def _realize_c4_pages(self, need_pages: int) -> None: base_backend admission 已按"空闲+可回收"放行本步请求,这里在真分配前(scatter 已算好 need) 把可回收的无引用 radix 节点驱逐出来腾出 c4 页,避免 alloc_c4_pages 触底 assert。 可回收仍不足时由 admission 的 wait_pause 兜底。""" - if self.n_c4 == 0 or need_pages <= 0: + if self.mem_manager.n_c4 == 0 or need_pages <= 0: return # 延迟 import: infer_batch 在模块顶 import 了 req_manager,顶层 import 会循环引用 from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -751,7 +719,7 @@ def _realize_c4_pages(self, need_pages: int) -> None: return def _realize_c128_slots(self, need_slots: int) -> None: - if self.n_c128 == 0 or need_slots <= 0: + if self.mem_manager.n_c128 == 0 or need_slots <= 0: return from lightllm.server.router.model_infer.infer_batch import g_infer_context @@ -768,12 +736,12 @@ def prepare_prefill_compress_slots( ) -> None: """prefill prep: 为本 chunk 内的组末 token(位置 (g+1)*ratio-1 ∈ [ready, seq))分配压缩槽, 组末 full 槽直接从 generic preprocess 的 mem_indexes 取。""" - if self.n_c4 == 0 and self.n_c128 == 0: + if self.mem_manager.n_c4 == 0 and self.mem_manager.n_c128 == 0: return - if self.n_c4 > 0: + if self.mem_manager.n_c4 > 0: self._scatter_c4_prefill_slots_batched(req_list, ready_list, seq_list, mem_indexes) - if self.n_c128 > 0: + if self.mem_manager.n_c128 > 0: ratio = 128 full_offsets = [] mem_offset = 0 @@ -805,9 +773,9 @@ def prepare_decode_compress_slots( """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 组末 full 槽即本步的 mem_index。 从 CPU 镜像读 seq_len/req_idx(host 算术,无 D2H);非关组步 rows 为空 => 不调 _scatter,零同步。""" - if self.n_c4 == 0 and self.n_c128 == 0: + if self.mem_manager.n_c4 == 0 and self.mem_manager.n_c128 == 0: return - if self.n_c4 > 0: + if self.mem_manager.n_c4 > 0: self._scatter_c4_decode_slots( req_list, seq_list, @@ -815,7 +783,7 @@ def prepare_decode_compress_slots( prev_group_end_mem_indexes=prev_group_end_mem_indexes, ) - if self.n_c128 > 0: + if self.mem_manager.n_c128 > 0: ratio = 128 rows = [ i diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index fd13bc43a0..e0640a52dd 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -67,15 +67,9 @@ def _init_req_manager(self): if self.max_seq_length is not None: create_max_seq_len = max(create_max_seq_len, self.max_seq_length) - self._dsv4_req_manager_seq_len = create_max_seq_len - layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() - self._dsv4_compress_rates = self._get_compress_rates(layer_num) self.req_manager = DeepseekV4ReqManager( self.max_req_num, create_max_seq_len, - compress_rates=self._dsv4_compress_rates, - head_dim=self.config["head_dim"], - indexer_head_dim=self.config["index_head_dim"], sliding_window=self.config["sliding_window"], ) return @@ -86,18 +80,15 @@ def _get_compress_rates(self, layer_num): def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() - compress_rates = getattr(self, "_dsv4_compress_rates", self._get_compress_rates(layer_num)) - sliding_window = int(self.config["sliding_window"]) self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, head_num=1, head_dim=self.config["head_dim"], layer_num=layer_num, - compress_rates=compress_rates, + compress_rates=self._get_compress_rates(layer_num), indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, - sliding_window=sliding_window, mem_fraction=self.mem_fraction, ) self.req_manager.mem_manager = self.mem_manager From dbe0f5e1844bb7b3d9b45f838cd656851b201f1d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 14 Jul 2026 11:42:16 +0000 Subject: [PATCH 080/214] reuse stable EP buffers to reduce memory fragmentation --- .../fused_moe/fused_moe_weight.py | 2 + .../fused_moe/impl/deepgemm_impl.py | 2 + .../fused_moe/impl/marlin_impl.py | 1 + .../meta_weights/fused_moe/impl/mxfp4_impl.py | 1 + .../fused_moe/impl/triton_impl.py | 3 + .../fused_moe/grouped_fused_moe_ep.py | 64 +- ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 74 ++ ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 632 ++++++++++++++++++ .../layer_infer/transformer_layer_infer.py | 1 + 9 files changed, 768 insertions(+), 12 deletions(-) create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json create mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 3e11e39852..bab48b9895 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -160,6 +160,7 @@ def experts_with_topk( topk_ids: torch.Tensor, is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ) -> torch.Tensor: return self.fuse_moe_impl.fused_experts_with_topk( input_tensor=input_tensor, @@ -169,6 +170,7 @@ def experts_with_topk( topk_ids=topk_ids, is_prefill=is_prefill, clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, ) def low_latency_dispatch( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 94709c5102..612c517155 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -77,6 +77,7 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ): output = fused_experts( hidden_states=input_tensor, @@ -89,6 +90,7 @@ def _fused_experts( is_prefill=is_prefill, previous_event=None, # for overlap clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, ) return output diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 1fdfd94d0d..a30a669c18 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -31,6 +31,7 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ): assert clamp_limit is None, "awq_marlin fused MoE does not support clamp_limit yet" diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py index 97cf238116..a7e19a9c80 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py @@ -19,6 +19,7 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ): try: from vllm.model_executor.layers.fused_moe.activation import MoEActivation diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index 8967dda34e..d8a3227236 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -95,6 +95,7 @@ def _fused_experts( router_logits: Optional[torch.Tensor] = None, is_prefill: bool = False, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -125,6 +126,7 @@ def fused_experts_with_topk( topk_ids: torch.Tensor, is_prefill: Optional[bool] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, ): return self._fused_experts( input_tensor=input_tensor, @@ -134,6 +136,7 @@ def fused_experts_with_topk( topk_ids=topk_ids, is_prefill=is_prefill, clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, ) def __call__( diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index c5b241f1ab..126e15326e 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -67,17 +67,20 @@ def masked_group_gemm( w2_scale: torch.Tensor, expected_m: int, clamp_limit: Optional[float] = None, + alloc_tensor_func: Callable = torch.empty, ): padded_m = recv_x[0].shape[1] E, N, _ = w1.shape block_size = 128 # groupgemm (masked layout) - gemm_out_a = torch.empty((E, padded_m, N), device=recv_x[0].device, dtype=dtype) + gemm_out_a = alloc_tensor_func((E, padded_m, N), device=recv_x[0].device, dtype=dtype) expected_m = min(expected_m, padded_m) - qsilu_out_scale = torch.empty((E, padded_m, N // 2 // block_size), device=recv_x[0].device, dtype=torch.float32) - qsilu_out = torch.empty((E, padded_m, N // 2), dtype=w1.dtype, device=recv_x[0].device) + qsilu_out_scale = alloc_tensor_func( + (E, padded_m, N // 2 // block_size), device=recv_x[0].device, dtype=torch.float32 + ) + qsilu_out = alloc_tensor_func((E, padded_m, N // 2), dtype=w1.dtype, device=recv_x[0].device) # groupgemm (masked layout) - gemm_out_b = torch.empty_like(recv_x[0], device=recv_x[0].device, dtype=dtype) + gemm_out_b = alloc_tensor_func(recv_x[0].shape, device=recv_x[0].device, dtype=dtype) _deepgemm_grouped_fp8_nt_masked(recv_x, (w1, w1_scale), gemm_out_a, masked_m, expected_m) @@ -121,6 +124,7 @@ def mega_moe_impl( topk_ids: torch.Tensor, quant_method: Any, clamp_limit: Optional[float] = None, + alloc_tensor_func: Callable = torch.empty, ): if not (HAS_DEEPGEMM and hasattr(deep_gemm, "fp8_fp4_mega_moe")): raise RuntimeError("deep_gemm does not provide fp8-fp4 Mega MoE kernel") @@ -151,7 +155,7 @@ def mega_moe_impl( buffer.topk_idx[:num_tokens].copy_(topk_ids) buffer.topk_weights[:num_tokens].copy_(topk_weights) - output = torch.empty_like(hidden_states) + output = alloc_tensor_func(hidden_states.shape, device=hidden_states.device, dtype=hidden_states.dtype) deep_gemm.fp8_fp4_mega_moe( output, l1_weights, @@ -197,10 +201,20 @@ def fused_experts( is_prefill: Optional[bool], previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func: Callable = torch.empty, ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): - return mega_moe_impl(hidden_states, w13, w2, topk_weights, topk_idx, quant_method, clamp_limit=clamp_limit) + return mega_moe_impl( + hidden_states, + w13, + w2, + topk_weights, + topk_idx, + quant_method, + clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, + ) buffer = dist_group_manager.ep_buffer if is_prefill else dist_group_manager.ep_low_latency_buffer return fused_experts_impl( @@ -219,6 +233,7 @@ def fused_experts( w2_scale=w2.weight_scale, previous_event=previous_event, clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, ) @@ -238,6 +253,7 @@ def fused_experts_impl( w2_scale: Optional[torch.Tensor] = None, previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, + alloc_tensor_func: Callable = torch.empty, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -262,7 +278,9 @@ def fused_experts_impl( combined_x = None if is_prefill: - qinput_tensor, input_scale = per_token_group_quant_fp8(hidden_states, block_size_k, dtype=w1.dtype) + qinput_tensor, input_scale = per_token_group_quant_fp8( + hidden_states, block_size_k, dtype=w1.dtype, alloc_func=alloc_tensor_func + ) allocate_on_comm_stream = previous_event is not None # normal dispatch # recv_x [recive_num_tokens, hidden] recv_x_scale [recive_num_tokens, hidden // block_size] @@ -304,7 +322,11 @@ def fused_experts_impl( handle.num_recv_tokens_per_expert_list, dtype=torch.int32, pin_memory=True, device="cpu" ).cuda(non_blocking=True) - expert_start_loc = torch.empty_like(num_recv_tokens_per_expert) + expert_start_loc = alloc_tensor_func( + num_recv_tokens_per_expert.shape, + device=num_recv_tokens_per_expert.device, + dtype=num_recv_tokens_per_expert.dtype, + ) ep_scatter( recv_x[0], @@ -317,11 +339,14 @@ def fused_experts_impl( m_indices, output_index, ) + recv_x = None # groupgemm (contiguous layout) - gemm_out_a = torch.empty((all_tokens, N), device=hidden_states.device, dtype=hidden_states.dtype) + gemm_storage = torch.empty(all_tokens * max(N, K), device=hidden_states.device, dtype=hidden_states.dtype) + gemm_out_a = gemm_storage[: all_tokens * N].view(all_tokens, N) input_tensor[1] = tma_align_input_scale(input_tensor[1]) deepgemm_grouped_fp8_nt_contiguous(input_tensor, (w1, w1_scale), gemm_out_a, m_indices) + input_tensor = None # silu_and_mul_fwd + qaunt # TODO fused kernel @@ -329,11 +354,17 @@ def fused_experts_impl( silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out, limit=clamp_limit) qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( - silu_out, block_size_k, dtype=w1.dtype, column_major_scales=True, scale_tma_aligned=True + silu_out, + block_size_k, + dtype=w1.dtype, + column_major_scales=True, + scale_tma_aligned=True, ) + gemm_out_a = None + silu_out = None # groupgemm (contiguous layout) - gemm_out_b = torch.empty((all_tokens, K), device=hidden_states.device, dtype=hidden_states.dtype) + gemm_out_b = gemm_storage[: all_tokens * K].view(all_tokens, K) deepgemm_grouped_fp8_nt_contiguous((qsilu_out, qsilu_out_scale), (w2, w2_scale), gemm_out_b, m_indices) @@ -372,7 +403,16 @@ def fused_experts_impl( ) # deepgemm gemm_out_b = masked_group_gemm( - recv_x, masked_m, hidden_states.dtype, w1, w1_scale, w2, w2_scale, expected_m, clamp_limit=clamp_limit + recv_x, + masked_m, + hidden_states.dtype, + w1, + w1_scale, + w2, + w2_scale, + expected_m, + clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, ) # low latency combine combined_x, event_overlap, hook = buffer.low_latency_combine( diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..e700378de1 --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,74 @@ +{ + "1": { + "BLOCK_SEQ": 8, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 4, + "num_warps": 4 + }, + "100": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 2 + }, + "1024": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 4, + "num_stages": 5, + "num_warps": 1 + }, + "128": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 2 + }, + "16": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 8 + }, + "2048": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 4, + "num_stages": 3, + "num_warps": 1 + }, + "256": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 8, + "num_stages": 2, + "num_warps": 2 + }, + "32": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 2 + }, + "4096": { + "BLOCK_SEQ": 2, + "HEAD_PARALLEL_NUM": 2, + "num_stages": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 2 + }, + "8": { + "BLOCK_SEQ": 1, + "HEAD_PARALLEL_NUM": 16, + "num_stages": 2, + "num_warps": 8 + }, + "8192": { + "BLOCK_SEQ": 8, + "HEAD_PARALLEL_NUM": 4, + "num_stages": 4, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json new file mode 100644 index 0000000000..fbd605f99a --- /dev/null +++ b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json @@ -0,0 +1,632 @@ +{ + "1": { + "BLOCK_M": 64, + "BLOCK_N": 32, + "NUM_STAGES": 2, + "num_warps": 8 + }, + "100": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "1024": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1152": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "12160": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "128": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1280": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "12928": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "13184": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13312": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13568": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13824": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "13952": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1408": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14080": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "14336": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "14592": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "14720": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "14848": { + "BLOCK_M": 64, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "14976": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "15232": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "1536": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "16": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "1664": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1792": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "1920": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2048": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2176": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2304": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "24192": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2432": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "24960": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25088": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25344": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "25472": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "256": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "2560": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25728": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "25856": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26112": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26240": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "26752": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2688": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27008": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27136": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27520": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27776": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "27904": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "28032": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2816": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "28288": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "28800": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "2944": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3072": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "32": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 4 + }, + "3200": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3328": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3456": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3584": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3712": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "384": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "3840": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "3968": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4096": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4224": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 1 + }, + "4352": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "4480": { + "BLOCK_M": 32, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "46336": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "46720": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "46976": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "48896": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "49152": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "49408": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50048": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50304": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50560": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50688": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50816": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "50944": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "512": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "51328": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "52608": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "53248": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "53632": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54144": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54272": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54528": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "54656": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "55040": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "64": { + "BLOCK_M": 1, + "BLOCK_N": 256, + "NUM_STAGES": 2, + "num_warps": 1 + }, + "640": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7424": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7552": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "768": { + "BLOCK_M": 8, + "BLOCK_N": 128, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7680": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7808": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "7936": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8": { + "BLOCK_M": 1, + "BLOCK_N": 128, + "NUM_STAGES": 2, + "num_warps": 4 + }, + "8064": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8192": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8320": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8448": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "8576": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + }, + "896": { + "BLOCK_M": 8, + "BLOCK_N": 256, + "NUM_STAGES": 1, + "num_warps": 4 + }, + "8960": { + "BLOCK_M": 32, + "BLOCK_N": 256, + "NUM_STAGES": 4, + "num_warps": 1 + } +} \ No newline at end of file diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index b7e125e547..638e11f6f2 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -358,6 +358,7 @@ def _routed_experts( topk_ids=indices, is_prefill=infer_state.is_prefill, clamp_limit=float(self.swiglu_limit), + alloc_tensor_func=self.alloc_tensor, ) def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): From 4ac53e82fd3ecabef8ff0bbd4100a5fec3202821 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 14 Jul 2026 12:02:45 +0000 Subject: [PATCH 081/214] add test/benchmark/static_inference/static_benchmark.py --- .../static_inference/static_benchmark.py | 1542 +++++++++++++++++ 1 file changed, 1542 insertions(+) create mode 100644 test/benchmark/static_inference/static_benchmark.py diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py new file mode 100644 index 0000000000..3fd3c62a3e --- /dev/null +++ b/test/benchmark/static_inference/static_benchmark.py @@ -0,0 +1,1542 @@ +"""Static forward benchmark for LightLLM model parts. + +The entry uses synthetic token ids and measures forward-only TPS for prefill, +chunked prefill, decode, and MTP decode cases. +""" + +import argparse +import json +import math +import os +import queue +import sys +import time +import traceback +from dataclasses import asdict, dataclass, replace +from pathlib import Path +from types import SimpleNamespace +from typing import Dict, List, Optional, Sequence + +import numpy as np +import torch +import torch.multiprocessing as mp +from transformers import PretrainedConfig + + +REPO_ROOT = Path(__file__).resolve().parents[3] +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.models import get_model +from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel +from lightllm.models.deepseek_v4_mtp.model import DeepseekV4MTPModel +from lightllm.models.glm4_moe_lite_mtp.model import Glm4MoeLiteMTPModel +from lightllm.models.mistral_mtp.model import MistralMTPModel +from lightllm.models.qwen3_moe_mtp.model import Qwen3MOEMTPModel +from lightllm.server.api_cli import make_argument_parser +from lightllm.utils.config_utils import get_dtype, get_vocab_size +from lightllm.utils.dist_utils import init_distributed_env +from lightllm.utils.envs_utils import set_env_start_args + + +DEFAULT_BATCH_SIZES = [2, 8, 16, 32, 64, 128] +MTP_MODES = {"vanilla_with_att", "eagle_with_att", "vanilla_no_att", "eagle_no_att"} +PREFILL_TABLE_HEADERS = [ + "ctx", + "hit", + "bs", + "max_total_token_num", + "uncached", + "cached", + "tokens", + "ms", + "qps", + "tok/s", + "logical_tok/s", +] +DECODE_TABLE_HEADERS = [ + "ctx", + "bs", + "accept", + "max_total_token_num", + "ms", + "qps", + "tok/s", + "itl_ms", +] + + +@dataclass(frozen=True) +class BenchmarkCase: + name: str + stage: str + batch_size: int + context_len: int + output_len: int + chunked_prefill_size: Optional[int] = None + profiled_max_total_token_num: Optional[int] = None + profiled_batch_divisor: Optional[int] = None + cache_hit_rate: float = 0.0 + prefill_uncached_len: Optional[int] = None + prefill_step_tokens_per_req: Optional[int] = None + prefill_batch_size_by_batch_max_tokens: Optional[int] = None + + +@dataclass +class BenchmarkResult: + case: str + stage: str + batch_size: int + context_len: int + output_len: int + chunked_prefill_size: Optional[int] + elapsed_ms: float + measured_tokens: int + qps: float + tps: float + profiled_max_total_token_num: Optional[int] = None + profiled_batch_divisor: Optional[int] = None + ttft_ms: Optional[float] = None + inter_token_latency_ms: Optional[float] = None + cache_hit_rate: float = 0.0 + prefill_uncached_len: Optional[int] = None + prefill_cached_len: Optional[int] = None + prefill_step_tokens_per_req: Optional[int] = None + mtp_accept_rate: Optional[float] = None + logical_tps: Optional[float] = None + + +class TokenSource: + def __init__(self, args: SimpleNamespace): + self.vocab_size = max(2, int(get_vocab_size(args.model_dir) or 0)) + self.rng = np.random.default_rng(args.seed) + + def batch(self, batch_size: int, need_len: int) -> np.ndarray: + return self.rng.integers(low=0, high=self.vocab_size, size=(batch_size, need_len), dtype=np.int64) + + +def cpu_i32_full(shape, value) -> torch.Tensor: + return torch.full(shape, value, dtype=torch.int32, device="cpu") + + +def cpu_i32_zeros(size: int) -> torch.Tensor: + return torch.zeros(size, dtype=torch.int32, device="cpu") + + +def empty_multimodal_params(batch_size: int) -> List[Dict]: + return [{"images": [], "audios": []} for _ in range(batch_size)] + + +class StaticBenchmarkExecutor: + def __init__( + self, + args: SimpleNamespace, + model, + draft_models: List, + token_source: TokenSource, + ): + self.args = args + self.model = model + self.draft_models = draft_models + self.token_source = token_source + + def _case_iters(self, warmup: bool) -> int: + return self.args.warmup_iters if warmup else self.args.bench_iters + + def run_case(self, case: BenchmarkCase, warmup: bool) -> BenchmarkResult: + if case.stage == "prefill": + return self._run_prefill_case(case, warmup) + if case.stage == "decode": + return self._run_decode_case(case, warmup) + raise ValueError(f"unknown benchmark stage: {case.stage}") + + def _run_prefill_case(self, case: BenchmarkCase, warmup: bool) -> BenchmarkResult: + """Measure full uncached prefill, chunked by production admission size.""" + uncached_len = int(case.prefill_uncached_len or case.context_len) + cached_len = max(0, case.context_len - uncached_len) + tokens = self.token_source.batch(case.batch_size, uncached_len) + elapsed = 0.0 + measured_tokens = case.batch_size * uncached_len + + for _ in range(self._case_iters(warmup)): + self._reset_model_cache() + req_idx = self._alloc_req_indexes(case.batch_size) + if cached_len > 0: + self._materialize_cached_prefix(req_idx, cached_len) + inputs = self._build_prefill_inputs( + token_rows=tokens, + req_idx=req_idx, + prompt_len=uncached_len, + chunk_size=case.chunked_prefill_size, + initial_ready_cache_len=cached_len, + ) + torch.cuda.synchronize() + start = time.perf_counter() + output = None + for model_input in inputs: + output = self._forward_prefill_input(model_input, allow_overlap=True) + torch.cuda.synchronize() + elapsed += time.perf_counter() - start + self._touch_output(output) + + self._reset_model_cache() + return self._make_result(case, elapsed, measured_tokens, warmup) + + def _run_decode_case(self, case: BenchmarkCase, warmup: bool) -> BenchmarkResult: + mtp_enabled = self._mtp_enabled() + measured_tokens = case.batch_size * case.output_len + elapsed = 0.0 + decode_step_count = 0 + iters = self._case_iters(warmup) + + for _ in range(iters): + self._reset_model_cache() + req_idx, seq_len, next_ids = self._materialize_context_for_decode(case) + if mtp_enabled: + step_elapsed, step_count = self._run_mtp_decode_steps( + case=case, + req_idx=req_idx, + seq_len=seq_len, + next_ids=next_ids, + ) + elapsed += step_elapsed + decode_step_count += step_count + else: + elapsed += self._run_plain_decode_steps( + case=case, + req_idx=req_idx, + seq_len=seq_len, + next_ids=next_ids, + ) + decode_step_count += case.output_len + + self._reset_model_cache() + inter_token_latency_ms = elapsed * 1000.0 / max(1, decode_step_count) if iters > 0 else None + return self._make_result( + case, + elapsed, + measured_tokens, + warmup, + inter_token_latency_ms=inter_token_latency_ms, + ) + + def _materialize_context_for_decode(self, case: BenchmarkCase): + """Allocate historical KV slots so decode can be measured without prefill.""" + req_idx = self._alloc_req_indexes(case.batch_size) + self._materialize_cached_prefix(req_idx, case.context_len) + seq_len = cpu_i32_full((case.batch_size,), case.context_len) + next_ids = torch.from_numpy(np.ascontiguousarray(self.token_source.batch(case.batch_size, 1).reshape(-1))).to( + torch.int64 + ) + return req_idx, seq_len, next_ids + + def _run_plain_decode_steps( + self, + case: BenchmarkCase, + req_idx: torch.Tensor, + seq_len: torch.Tensor, + next_ids: torch.Tensor, + ) -> float: + elapsed = 0.0 + for step in range(case.output_len): + seq_len += 1 + model_input = self._make_decode_input( + batch_size=case.batch_size, + req_idx=req_idx, + mtp_index=cpu_i32_zeros(case.batch_size), + seq_len=seq_len, + input_ids=next_ids.reshape(-1), + max_kv_seq_len=int(seq_len.max().item()), + mem_token_num=case.batch_size, + ) + torch.cuda.synchronize() + start = time.perf_counter() + output = self._forward_decode_input(model_input, allow_overlap=True) + torch.cuda.synchronize() + elapsed += time.perf_counter() - start + self._touch_output(output) + + next_ids = self._argmax_ids(output.logits) + return elapsed + + def _run_mtp_decode_steps( + self, + case: BenchmarkCase, + req_idx: torch.Tensor, + seq_len: torch.Tensor, + next_ids: torch.Tensor, + ) -> tuple: + elapsed = 0.0 + step_count = 0 + generated_len = 0 + step_width = self._mtp_step_width() + base_req_idx, b_mtp_index = self._build_mtp_decode_index_tensors(req_idx, step_width) + current_candidates = next_ids + + while generated_len < case.output_len: + accepted_width = self._sample_mtp_accept_width(step_width, case.output_len - generated_len) + if current_candidates.ndim == 1: + current_candidates = current_candidates[:, None].repeat(1, step_width) + + b_seq_len = self._build_mtp_seq_len(seq_len, step_width) + model_input = self._make_decode_input( + batch_size=case.batch_size * step_width, + req_idx=base_req_idx, + mtp_index=b_mtp_index, + seq_len=b_seq_len, + input_ids=current_candidates.reshape(-1), + max_kv_seq_len=int(b_seq_len.max().item()), + mem_token_num=case.batch_size * step_width, + ) + + torch.cuda.synchronize() + start = time.perf_counter() + output = self.model.forward(model_input) + candidate_rows, temporary_mem = self._run_mtp_draft_decode( + model_input=model_input, + model_output=output, + real_batch_size=case.batch_size, + step_width=step_width, + ) + torch.cuda.synchronize() + elapsed += time.perf_counter() - start + self._touch_output(output) + if temporary_mem is not None: + self.model.req_manager.mem_manager.free(temporary_mem) + + self._free_rejected_mtp_mem( + model_input=model_input, + real_batch_size=case.batch_size, + step_width=step_width, + accepted_width=accepted_width, + ) + current_candidates = ( + self._select_mtp_candidates( + candidate_rows=candidate_rows, + real_batch_size=case.batch_size, + step_width=step_width, + accepted_width=accepted_width, + ) + .detach() + .cpu() + ) + seq_len += accepted_width + generated_len += accepted_width + step_count += 1 + + return elapsed, step_count + + def _run_mtp_draft_decode( + self, + model_input: ModelInput, + model_output: ModelOutput, + real_batch_size: int, + step_width: int, + ): + draft_input = model_input.make_mtp_draft_input() + draft_output = model_output + draft_next_ids = self._argmax_ids(model_output.logits).cuda(non_blocking=True) + generated = [draft_next_ids.detach()] + + temporary_mem = None + if self.args.mtp_mode.startswith("eagle"): + temporary_mem = self.model.req_manager.mem_manager.alloc(real_batch_size * self.args.mtp_step) + temporary_mem_gpu = temporary_mem.cuda(non_blocking=True) + else: + temporary_mem_gpu = None + + for step in range(self.args.mtp_step): + draft_input.input_ids = draft_next_ids + draft_input.mtp_draft_input_hiddens = draft_output.mtp_main_output_hiddens + draft_model = self.draft_models[step % self._num_mtp_modules()] + draft_output = draft_model.forward(draft_input) + draft_next_ids = self._argmax_ids(draft_output.logits).cuda(non_blocking=True) + generated.append(draft_next_ids.detach()) + + if self.args.mtp_mode.startswith("eagle"): + mem_i_cpu = temporary_mem[step * real_batch_size : (step + 1) * real_batch_size] + mem_i = temporary_mem_gpu[step * real_batch_size : (step + 1) * real_batch_size] + draft_input.advance_mtp_decode_step(mem_i_cpu, mem_i, self.args.mtp_step) + + return torch.stack(generated[:step_width], dim=1), temporary_mem + + def _sample_mtp_accept_width(self, step_width: int, remaining_tokens: int) -> int: + """Sample accepted MTP width outside the timed decode section.""" + accept_rate = float(self.args.mtp_accept_rate) + accepted_width = 1 + for _ in range(step_width - 1): + if self.token_source.rng.random() >= accept_rate: + break + accepted_width += 1 + return max(1, min(accepted_width, remaining_tokens)) + + def _select_mtp_candidates( + self, + candidate_rows: torch.Tensor, + real_batch_size: int, + step_width: int, + accepted_width: int, + ) -> torch.Tensor: + row_ids = torch.arange(real_batch_size, device=candidate_rows.device) * step_width + accepted_width - 1 + return candidate_rows.index_select(0, row_ids) + + def _free_rejected_mtp_mem( + self, + model_input: ModelInput, + real_batch_size: int, + step_width: int, + accepted_width: int, + ): + if accepted_width >= step_width: + return + rejected_mem = ( + model_input.mem_indexes_cpu.view(real_batch_size, step_width)[:, accepted_width:].contiguous().reshape(-1) + ) + if rejected_mem.numel() > 0: + self.model.req_manager.mem_manager.free(rejected_mem) + + def _build_prefill_inputs( + self, + token_rows: np.ndarray, + req_idx: torch.Tensor, + prompt_len: int, + chunk_size: Optional[int], + initial_ready_cache_len: int = 0, + ) -> List[ModelInput]: + if not chunk_size or chunk_size <= 0 or chunk_size >= prompt_len: + return [ + self._make_prefill_input( + token_rows[:, :prompt_len], + req_idx, + ready_cache_len=initial_ready_cache_len, + ) + ] + + inputs = [] + for start in range(0, prompt_len, chunk_size): + end = min(prompt_len, start + chunk_size) + inputs.append( + self._make_prefill_input( + token_rows[:, start:end], + req_idx, + ready_cache_len=initial_ready_cache_len + start, + ) + ) + return inputs + + def _materialize_cached_prefix(self, req_idx: torch.Tensor, cached_len: int): + """Allocate dummy prefix KV so cache-hit cases consume real capacity.""" + if cached_len <= 0: + return + batch_size = int(req_idx.shape[0]) + need_tokens = batch_size * cached_len + mem_indexes = self.model.req_manager.mem_manager.alloc(need_tokens) + if mem_indexes is None: + raise RuntimeError(f"failed to allocate cached prefix: bs={batch_size} cached_len={cached_len}") + req_idx_gpu = req_idx.cuda(non_blocking=True) + mem_indexes_gpu = mem_indexes.reshape(batch_size, cached_len).cuda(non_blocking=True) + self.model.req_manager.req_to_token_indexs[req_idx_gpu, :cached_len] = mem_indexes_gpu + self._materialize_cached_prefix_extra_slots(req_idx, cached_len) + + def _materialize_cached_prefix_extra_slots(self, req_idx: torch.Tensor, cached_len: int): + req_manager = self.model.req_manager + batch_size = int(req_idx.shape[0]) + b_req_idx = req_idx.cuda(non_blocking=True) + b_seq_len_cpu = cpu_i32_full((batch_size,), cached_len) + b_seq_len = b_seq_len_cpu.cuda(non_blocking=True) + + if hasattr(req_manager, "prepare_prefill_compress_slots"): + b_ready_cache_len_cpu = cpu_i32_zeros(batch_size) + req_manager.prepare_prefill_compress_slots( + b_req_idx=b_req_idx, + b_ready_cache_len=b_ready_cache_len_cpu.cuda(non_blocking=True), + b_seq_len=b_seq_len, + b_req_idx_cpu=req_idx, + b_ready_cache_len_cpu=b_ready_cache_len_cpu, + b_seq_len_cpu=b_seq_len_cpu, + ) + + if hasattr(req_manager, "prepare_prefill_swa"): + swa_ready_len = self._cached_prefix_swa_ready_len(cached_len) + b_ready_cache_len_cpu = cpu_i32_full((batch_size,), swa_ready_len) + req_manager.prepare_prefill_swa( + b_req_idx=b_req_idx, + b_ready_cache_len=b_ready_cache_len_cpu.cuda(non_blocking=True), + b_seq_len=b_seq_len, + b_req_idx_cpu=req_idx, + b_ready_cache_len_cpu=b_ready_cache_len_cpu, + b_seq_len_cpu=b_seq_len_cpu, + ) + + def _cached_prefix_swa_ready_len(self, cached_len: int) -> int: + req_manager = self.model.req_manager + retain_fn = getattr(req_manager, "_swa_retain_len", None) + if retain_fn is not None: + retain_len = int(retain_fn()) + else: + retain_len = int(getattr(req_manager, "sliding_window", cached_len) or cached_len) + ready_len = max(0, int(cached_len) - max(1, retain_len)) + page_fn = getattr(req_manager, "get_prompt_cache_page_size", None) + page_size = int(page_fn()) if page_fn is not None else 1 + return ready_len // max(1, page_size) * max(1, page_size) + + def _make_prefill_input(self, token_chunk: np.ndarray, req_idx: torch.Tensor, ready_cache_len: int) -> ModelInput: + batch_size, q_len = token_chunk.shape + seq_len_value = ready_cache_len + q_len + b_seq_len = cpu_i32_full((batch_size,), seq_len_value) + b_ready_cache_len = cpu_i32_full((batch_size,), ready_cache_len) + 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 + input_ids = torch.from_numpy(np.ascontiguousarray(token_chunk.reshape(-1))).to(torch.int64) + mem_indexes = self.model.req_manager.mem_manager.alloc(input_ids.shape[0]) + return ModelInput( + batch_size=batch_size, + total_token_num=int(b_seq_len.sum().item()), + 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), + b_seq_len=b_seq_len, + mem_indexes_cpu=mem_indexes, + 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] * batch_size, + multimodal_params=empty_multimodal_params(batch_size), + ) + + def _make_decode_input( + self, + batch_size: int, + req_idx: torch.Tensor, + mtp_index: torch.Tensor, + seq_len: torch.Tensor, + input_ids: torch.Tensor, + max_kv_seq_len: int, + mem_token_num: int, + ) -> ModelInput: + mem_indexes = self.model.req_manager.mem_manager.alloc(mem_token_num) + return ModelInput( + batch_size=batch_size, + total_token_num=int(seq_len.sum().item()), + max_q_seq_len=1, + max_kv_seq_len=max_kv_seq_len, + input_ids=input_ids.to(torch.int64).cpu(), + b_req_idx=req_idx, + b_mtp_index=mtp_index, + b_seq_len=seq_len, + mem_indexes_cpu=mem_indexes, + is_prefill=False, + multimodal_params=empty_multimodal_params(batch_size), + ) + + 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) + 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) + return self.model.forward(model_input) + + def _split_prefill_input(self, model_input: ModelInput): + 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()) + return ( + self._slice_prefill_input(model_input, 0, split_batch, 0, split_tokens), + self._slice_prefill_input( + model_input, + split_batch, + model_input.batch_size, + split_tokens, + int(q_lens.sum().item()), + ), + ) + + def _slice_prefill_input( + self, + model_input: ModelInput, + batch_start: int, + batch_end: int, + token_start: int, + token_end: int, + ) -> 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 + has_output = model_input.b_prefill_has_output_cpu + return ModelInput( + 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_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(), + 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), + multimodal_params=model_input.multimodal_params[batch_start:batch_end], + ) + + def _split_decode_input(self, model_input: ModelInput): + split_batch = model_input.batch_size // 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: + 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, + 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()), + 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, + mem_indexes_cpu=model_input.mem_indexes_cpu[batch_start:batch_end].contiguous(), + is_prefill=False, + multimodal_params=model_input.multimodal_params[batch_start:batch_end], + ) + + 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, + ) + + def _build_mtp_decode_index_tensors(self, req_idx: torch.Tensor, step_width: int): + batch_size = int(req_idx.shape[0]) + return ( + req_idx.repeat_interleave(step_width).to(torch.int32).cpu(), + torch.arange(step_width, dtype=torch.int32).repeat(batch_size), + ) + + def _build_mtp_seq_len(self, base_seq_len: torch.Tensor, step_width: int) -> torch.Tensor: + offsets = torch.arange(1, step_width + 1, dtype=torch.int32) + return (base_seq_len[:, None].to(torch.int32) + offsets[None, :]).reshape(-1) + + def _alloc_req_indexes(self, batch_size: int) -> torch.Tensor: + req_indexes = [self.model.req_manager.alloc() for _ in range(batch_size)] + if any(index is None for index in req_indexes): + raise RuntimeError(f"failed to allocate {batch_size} request indexes") + return torch.tensor(req_indexes, dtype=torch.int32, device="cpu") + + def _reset_model_cache(self): + self.model.mem_manager.free_all() + self.model.req_manager.free_all() + torch.cuda.synchronize() + torch.cuda.empty_cache() + + def _argmax_ids(self, logits: torch.Tensor) -> torch.Tensor: + return torch.argmax(logits, dim=-1).detach().cpu().to(torch.int64) + + def _touch_output(self, output: Optional[ModelOutput]): + if output is not None and output.logits is not None: + _ = output.logits.shape + + def _make_result( + self, + case: BenchmarkCase, + elapsed_s: float, + measured_tokens: int, + warmup: bool, + ttft_elapsed_s: Optional[float] = None, + inter_token_latency_ms: Optional[float] = None, + ) -> BenchmarkResult: + """Convert raw timings into reported TPS and latency metrics.""" + iters = self._case_iters(warmup) + scaled_tokens = measured_tokens * iters + qps = case.batch_size * iters / elapsed_s if elapsed_s > 0 else 0.0 + tps = scaled_tokens / elapsed_s if elapsed_s > 0 else 0.0 + ttft_ms = ttft_elapsed_s * 1000.0 / max(1, iters) if ttft_elapsed_s is not None else None + logical_tps = None + prefill_uncached_len = case.prefill_uncached_len + prefill_cached_len = None + if case.stage == "prefill": + uncached_len = int(case.prefill_uncached_len or case.context_len) + prefill_uncached_len = uncached_len + prefill_cached_len = max(0, case.context_len - uncached_len) + token_count = case.batch_size * case.context_len * iters + logical_tps = token_count / elapsed_s if elapsed_s > 0 else 0.0 + return BenchmarkResult( + case=case.name, + stage=case.stage, + batch_size=case.batch_size, + context_len=case.context_len, + output_len=case.output_len, + chunked_prefill_size=case.chunked_prefill_size, + elapsed_ms=elapsed_s * 1000.0, + measured_tokens=scaled_tokens, + qps=qps, + tps=tps, + profiled_max_total_token_num=case.profiled_max_total_token_num, + profiled_batch_divisor=case.profiled_batch_divisor, + ttft_ms=ttft_ms, + inter_token_latency_ms=inter_token_latency_ms, + cache_hit_rate=case.cache_hit_rate, + prefill_uncached_len=prefill_uncached_len, + prefill_cached_len=prefill_cached_len, + prefill_step_tokens_per_req=case.prefill_step_tokens_per_req, + mtp_accept_rate=( + float(self.args.mtp_accept_rate) if case.stage == "decode" and self._mtp_enabled() else None + ), + logical_tps=logical_tps, + ) + + def _mtp_enabled(self) -> bool: + return self.args.mtp_mode in MTP_MODES and self.args.mtp_step > 0 + + def _mtp_step_width(self) -> int: + return int(self.args.mtp_step) + 1 + + def _num_mtp_modules(self) -> int: + if not self._mtp_enabled(): + return 0 + if self.args.mtp_mode.startswith("eagle"): + return 1 + return int(self.args.mtp_step) + + +def parse_typed_list(value: Optional[str], fallback: Sequence, cast) -> List: + if value is None or value == "": + return list(fallback) + if isinstance(value, cast): + return [value] + normalized = str(value).replace(",", " ") + return [cast(item) for item in normalized.split() if item.strip()] + + +def parse_int_list(value: Optional[str], fallback: Sequence[int]) -> List[int]: + return parse_typed_list(value, fallback, int) + + +def parse_float_list(value: Optional[str], fallback: Sequence[float]) -> List[float]: + return parse_typed_list(value, fallback, float) + + +def parse_chunk_sizes(value: Optional[str], fallback: Optional[int]) -> List[Optional[int]]: + if value is None: + return [fallback] if fallback else [None] + chunks: List[Optional[int]] = [] + for item in str(value).replace(",", " ").split(): + item = item.strip().lower() + if item in {"none", "full", "0", "-1"}: + chunks.append(None) + else: + chunks.append(int(item)) + return chunks or [None] + + +def prefill_uncached_len(context_len: int, cache_hit_rate: float) -> int: + """Return the uncached suffix length for a prompt-cache hit ratio.""" + if cache_hit_rate < 0.0 or cache_hit_rate >= 1.0: + raise ValueError(f"cache hit rate must satisfy 0 <= hit < 1, got {cache_hit_rate}") + uncached = int(math.ceil(context_len * (1.0 - cache_hit_rate))) + return max(1, min(context_len, uncached)) + + +def prefill_step_tokens_per_req(uncached_len: int, chunked_prefill_size: Optional[int]) -> int: + """Return tokens handled per request in one production prefill step.""" + if chunked_prefill_size and chunked_prefill_size > 0: + return max(1, min(uncached_len, int(chunked_prefill_size))) + return max(1, uncached_len) + + +def format_cache_hit_suffix(cache_hit_rate: float) -> str: + return f"{cache_hit_rate:.4f}".rstrip("0").rstrip(".").replace(".", "p") + + +def apply_max_batch_size(batch_size: int, max_batch_size: int) -> int: + """Apply the benchmark-wide auto batch-size upper bound.""" + if max_batch_size > 0: + batch_size = min(batch_size, int(max_batch_size)) + return max(1, batch_size) + + +def prefill_batch_size_from_batch_max_tokens( + batch_max_tokens: int, + step_tokens_per_req: int, + max_batch_size: int, +) -> int: + """Compute prefill BS from batch_max_tokens before KV-capacity capping.""" + batch_size = max(1, int(batch_max_tokens) // max(1, step_tokens_per_req)) + return apply_max_batch_size(batch_size, max_batch_size) + + +def build_prefill_cases( + args: SimpleNamespace, + input_lens: Sequence[int], + chunk_sizes: Sequence[Optional[int]], + cache_hit_rates: Sequence[float], +) -> List[BenchmarkCase]: + """Build full-prefill cases using batch_max_tokens per chunk step.""" + if args.batch_max_tokens is None: + raise ValueError("prefill benchmark requires --batch_max_tokens") + cases: List[BenchmarkCase] = [] + for input_len in input_lens: + for chunk_size in chunk_sizes: + for cache_hit_rate in cache_hit_rates: + uncached_len = prefill_uncached_len(input_len, cache_hit_rate) + step_tokens = prefill_step_tokens_per_req(uncached_len, chunk_size) + bs = prefill_batch_size_from_batch_max_tokens( + args.batch_max_tokens, + step_tokens, + args.max_batch_size, + ) + chunk_name = chunk_size if chunk_size else "none" + hit_name = format_cache_hit_suffix(cache_hit_rate) + cases.append( + BenchmarkCase( + name=( + f"prefill_bs{bs}_in{input_len}_hit{hit_name}" + f"_uncached{uncached_len}_chunk{chunk_name}" + f"_btok{args.batch_max_tokens}" + ), + stage="prefill", + batch_size=bs, + context_len=input_len, + output_len=0, + chunked_prefill_size=chunk_size, + cache_hit_rate=cache_hit_rate, + prefill_uncached_len=uncached_len, + prefill_step_tokens_per_req=step_tokens, + prefill_batch_size_by_batch_max_tokens=bs, + ) + ) + return cases + + +def build_decode_cases( + args: SimpleNamespace, + batch_sizes: Sequence[int], + context_lens: Sequence[int], + output_lens: Sequence[int], +) -> List[BenchmarkCase]: + """Build decode cases; profile mode resolves BS after model load.""" + decode_batch_sizes = [1] if args.decode_batch_size_mode == "profile" else batch_sizes + cases: List[BenchmarkCase] = [] + for bs in decode_batch_sizes: + for context_len in context_lens: + for output_len in output_lens: + profile_suffix = "_profilebs" if args.decode_batch_size_mode == "profile" else "" + cases.append( + BenchmarkCase( + name=(f"decode_bs{bs}_ctx{context_len}_out{output_len}" f"{profile_suffix}"), + stage="decode", + batch_size=bs, + context_len=context_len, + output_len=output_len, + ) + ) + return cases + + +def build_cases(args: SimpleNamespace) -> List[BenchmarkCase]: + """Expand CLI list options into concrete prefill/decode benchmark cases.""" + batch_sizes = parse_int_list(args.batch_sizes, [args.batch_size] if args.batch_size else DEFAULT_BATCH_SIZES) + input_lens = parse_int_list(args.input_lens, [args.input_len]) + context_lens = parse_int_list(args.context_lens, input_lens) + output_lens = parse_int_list(args.output_lens, [args.output_len]) + chunk_sizes = parse_chunk_sizes(args.chunked_prefill_sizes, args.chunked_prefill_size) + cache_hit_rates = parse_float_list(args.prefill_cache_hit_rates, [0.0]) + + cases: List[BenchmarkCase] = [] + if args.benchmark in {"all", "prefill"}: + cases.extend(build_prefill_cases(args, input_lens, chunk_sizes, cache_hit_rates)) + if args.benchmark in {"all", "decode"}: + cases.extend(build_decode_cases(args, batch_sizes, context_lens, output_lens)) + return cases + + +def decode_profile_batch_divisor(args: SimpleNamespace, case: BenchmarkCase) -> int: + """Reserve KV capacity for context, generated tokens, and MTP expansion.""" + mtp_width = int(args.mtp_step) + 1 if args.mtp_mode in MTP_MODES else 1 + return max(1, case.context_len + case.output_len + mtp_width + 8) + + +def filter_capacity_decode_cases( + args: SimpleNamespace, + cases: Sequence[BenchmarkCase], + profiled_max_total_token_num: int, +) -> List[BenchmarkCase]: + if not getattr(args, "decode_filter_capacity", False): + return list(cases) + + resolved: List[BenchmarkCase] = [] + capacity_tokens = int(profiled_max_total_token_num) + for case in cases: + if case.stage != "decode": + resolved.append(case) + continue + divisor = decode_profile_batch_divisor(args, case) + if case.batch_size * divisor <= capacity_tokens: + resolved.append( + replace( + case, + profiled_max_total_token_num=capacity_tokens, + profiled_batch_divisor=divisor, + ) + ) + return resolved + + +def resolve_profile_decode_cases( + args: SimpleNamespace, + cases: Sequence[BenchmarkCase], + profiled_max_total_token_num: int, +) -> List[BenchmarkCase]: + """Replace decode profile placeholders with capacity-derived max BS.""" + if args.decode_batch_size_mode != "profile": + return list(cases) + + resolved: List[BenchmarkCase] = [] + for case in cases: + if case.stage != "decode": + resolved.append(case) + continue + + divisor = decode_profile_batch_divisor(args, case) + batch_size = max(1, int(profiled_max_total_token_num) // divisor) + batch_size = apply_max_batch_size(batch_size, args.max_batch_size) + + resolved.append( + replace( + case, + name=( + f"decode_bs{batch_size}_ctx{case.context_len}" + f"_out{case.output_len}_profile{profiled_max_total_token_num}" + ), + batch_size=batch_size, + profiled_max_total_token_num=int(profiled_max_total_token_num), + profiled_batch_divisor=divisor, + ) + ) + + return resolved + + +def resolve_batch_max_prefill_cases( + args: SimpleNamespace, + cases: Sequence[BenchmarkCase], + profiled_max_total_token_num: int, +) -> List[BenchmarkCase]: + """Cap prefill BS by profiled KV capacity after the model is loaded.""" + resolved: List[BenchmarkCase] = [] + capacity_tokens = int(profiled_max_total_token_num) + for case in cases: + if case.stage != "prefill": + resolved.append(case) + continue + + if case.context_len <= 0: + raise ValueError(f"invalid prefill context_len={case.context_len}") + bs_by_capacity = capacity_tokens // case.context_len + if bs_by_capacity <= 0: + raise ValueError( + "single prefill request does not fit profiled token capacity: " + f"context_len={case.context_len} capacity={profiled_max_total_token_num}" + ) + + bs_by_batch = int(case.prefill_batch_size_by_batch_max_tokens or case.batch_size) + batch_size = min(bs_by_batch, bs_by_capacity) + batch_size = apply_max_batch_size(batch_size, args.max_batch_size) + + chunk_name = case.chunked_prefill_size if case.chunked_prefill_size else "none" + hit_name = format_cache_hit_suffix(case.cache_hit_rate) + resolved.append( + replace( + case, + name=( + f"prefill_bs{batch_size}_in{case.context_len}_hit{hit_name}" + f"_uncached{case.prefill_uncached_len}_chunk{chunk_name}" + f"_btok{args.batch_max_tokens}" + f"_cap{profiled_max_total_token_num}" + ), + batch_size=batch_size, + profiled_max_total_token_num=int(profiled_max_total_token_num), + profiled_batch_divisor=case.context_len, + ) + ) + + return resolved + + +def normalize_args(args: argparse.Namespace, cases: Sequence[BenchmarkCase]) -> SimpleNamespace: + """Fill LightLLM startup args needed before model construction.""" + if args.data_type is None: + args.data_type = get_dtype(args.model_dir) + + if args.quant_type is None: + args.quant_type = "none" + + if not 0.0 <= float(args.mtp_accept_rate) <= 1.0: + raise ValueError(f"--mtp_accept_rate must be in [0, 1], got {args.mtp_accept_rate}") + + max_batch = max(case.batch_size for case in cases) + max_context = max(case.context_len for case in cases) + max_output = max(case.output_len for case in cases) + mtp_width = (args.mtp_step + 1) if args.mtp_mode in MTP_MODES else 1 + max_runtime_len = max_context + max_output + mtp_width + 2 + + if args.max_req_total_len is None: + args.max_req_total_len = max_runtime_len + else: + args.max_req_total_len = max(args.max_req_total_len, max_runtime_len) + + if args.graph_max_len_in_batch == 0: + args.graph_max_len_in_batch = args.max_req_total_len + + max_prefill_chunk = ( + max( + min(case.context_len, case.chunked_prefill_size or case.context_len) + for case in cases + if case.stage == "prefill" + ) + if any(case.stage == "prefill" for case in cases) + else max_context + ) + if args.batch_max_tokens is None: + args.batch_max_tokens = max(max_batch * max_prefill_chunk, max_batch * mtp_width, 1) + + decode_batch_size_needs_profile = ( + args.benchmark in {"all", "decode"} + and args.decode_batch_size_mode == "profile" + and args.max_total_token_num is None + ) + prefill_batch_size_needs_profile = args.benchmark in {"all", "prefill"} and args.max_total_token_num is None + needs_profiled_batch_size = decode_batch_size_needs_profile or prefill_batch_size_needs_profile + + if args.max_total_token_num is None and not needs_profiled_batch_size: + args.max_total_token_num = max_batch * (args.max_req_total_len + mtp_width + 8) + if args.max_total_token_num is not None: + args.max_total_token_num = max(args.max_total_token_num, args.batch_max_tokens + 1, args.max_req_total_len) + + if decode_batch_size_needs_profile and args.max_batch_size > 0: + args.running_max_req_size = max(args.running_max_req_size, int(args.max_batch_size)) + # Profile decode BS is resolved after model load. Use the cap as the + # pre-load upper bound so request slots and optional decode graphs agree. + if not args.disable_cudagraph: + args.graph_max_batch_size = max(args.graph_max_batch_size, int(args.max_batch_size)) + if prefill_batch_size_needs_profile: + args.running_max_req_size = max(args.running_max_req_size, max_batch) + + if args.graph_max_batch_size < max_batch: + args.graph_max_batch_size = max_batch + + if args.nccl_port is None: + args.nccl_port = 28765 + + if args.mtp_mode in MTP_MODES: + if args.mtp_step <= 0: + raise ValueError("--mtp_mode requires --mtp_step > 0") + if not args.mtp_draft_model_dir: + raise ValueError("--mtp_mode requires --mtp_draft_model_dir") + args.mtp_draft_model_dir = normalize_mtp_draft_dirs(args.mtp_mode, args.mtp_step, args.mtp_draft_model_dir) + else: + args.mtp_mode = None + args.mtp_step = 0 + args.mtp_draft_model_dir = None + + return SimpleNamespace(**vars(args)) + + +def normalize_mtp_draft_dirs(mtp_mode: str, mtp_step: int, draft_dirs: Sequence[str]) -> List[str]: + expected = 1 if mtp_mode.startswith("eagle") else mtp_step + if isinstance(draft_dirs, str): + draft_dirs = [draft_dirs] + draft_dirs = list(draft_dirs) + if len(draft_dirs) == 1 and expected > 1: + return draft_dirs * expected + if len(draft_dirs) != expected: + raise ValueError(f"{mtp_mode} expects {expected} draft model dir(s), got {len(draft_dirs)}") + return draft_dirs + + +def build_model_kvargs(args: SimpleNamespace, rank_id: int) -> Dict: + return { + "args": args, + "nccl_host": args.nccl_host, + "nccl_port": args.nccl_port, + "rank_id": rank_id, + "world_size": args.tp, + "dp_size": args.dp, + "weight_dir": args.model_dir, + "data_type": args.data_type, + "quant_type": args.quant_type, + "quant_cfg": args.quant_cfg, + "expert_dtype": args.expert_dtype, + "load_way": "HF", + "max_total_token_num": args.max_total_token_num, + "graph_max_len_in_batch": args.graph_max_len_in_batch, + "graph_max_batch_size": args.graph_max_batch_size, + "mem_fraction": args.mem_fraction, + "max_req_num": max(args.running_max_req_size, args.graph_max_batch_size), + "batch_max_tokens": args.batch_max_tokens, + "run_mode": "normal", + "max_seq_length": args.max_req_total_len, + "disable_cudagraph": args.disable_cudagraph, + "llm_prefill_att_backend": args.llm_prefill_att_backend, + "llm_decode_att_backend": args.llm_decode_att_backend, + "vit_att_backend": args.vit_att_backend, + "llm_kv_type": args.llm_kv_type, + "llm_kv_quant_group_size": args.llm_kv_quant_group_size, + } + + +def init_mtp_draft_models(args: SimpleNamespace, main_kvargs: Dict, main_model) -> List: + if args.mtp_mode not in MTP_MODES: + return [] + + os.environ["DISABLE_CHECK_MAX_LEN_INFER"] = "1" + draft_models = [] + for draft_dir in args.mtp_draft_model_dir: + mtp_cfg, _ = PretrainedConfig.get_config_dict(draft_dir) + model_type = mtp_cfg.get("model_type", "") + mtp_kvargs = { + "weight_dir": draft_dir, + "max_total_token_num": main_model.mem_manager.size, + "load_way": main_kvargs["load_way"], + "max_req_num": main_kvargs["max_req_num"], + "max_seq_length": main_kvargs["max_seq_length"], + "is_token_healing": False, + "return_all_prompt_logics": False, + "disable_chunked_prefill": args.disable_chunked_prefill, + "data_type": main_kvargs["data_type"], + "graph_max_batch_size": main_kvargs["graph_max_batch_size"], + "graph_max_len_in_batch": main_kvargs["graph_max_len_in_batch"], + "disable_cudagraph": main_kvargs["disable_cudagraph"], + "mem_fraction": main_kvargs["mem_fraction"], + "batch_max_tokens": main_kvargs["batch_max_tokens"], + "quant_type": main_kvargs["quant_type"], + "quant_cfg": main_kvargs["quant_cfg"], + "expert_dtype": main_kvargs["expert_dtype"], + "run_mode": "normal", + "main_model": main_model, + "mtp_previous_draft_models": draft_models.copy(), + } + if model_type == "deepseek_v3": + assert args.mtp_mode in { + "vanilla_with_att", + "eagle_with_att", + }, f"{model_type} MTP requires *_with_att mode" + draft_models.append(Deepseek3MTPModel(mtp_kvargs)) + elif model_type == "deepseek_v4": + assert args.mtp_mode == "eagle_with_att", f"{model_type} MTP requires eagle_with_att mode" + draft_models.append(DeepseekV4MTPModel(mtp_kvargs)) + elif model_type == "qwen3_moe": + assert args.mtp_mode in { + "vanilla_no_att", + "eagle_no_att", + }, f"{model_type} MTP requires *_no_att mode" + draft_models.append(Qwen3MOEMTPModel(mtp_kvargs)) + elif model_type == "mistral": + assert args.mtp_mode in { + "vanilla_no_att", + "eagle_no_att", + }, f"{model_type} MTP requires *_no_att mode" + draft_models.append(MistralMTPModel(mtp_kvargs)) + elif model_type == "glm4_moe_lite": + assert args.mtp_mode in { + "vanilla_with_att", + "eagle_with_att", + }, f"{model_type} MTP requires *_with_att mode" + draft_models.append(Glm4MoeLiteMTPModel(mtp_kvargs)) + else: + raise ValueError(f"unsupported MTP draft model_type={model_type} from {draft_dir}") + return draft_models + + +def run_worker(args_dict: Dict, case_dicts: List[Dict], rank_id: int, ans_queue): + try: + args = SimpleNamespace(**args_dict) + cases = [BenchmarkCase(**case) for case in case_dicts] + set_env_start_args(args) + + from lightllm.distributed import dist_group_manager + import torch.distributed as dist + + model_kvargs = build_model_kvargs(args, rank_id) + group_size = 2 if (args.enable_decode_microbatch_overlap or args.enable_prefill_microbatch_overlap) else 1 + if group_size == 2: + for case in cases: + assert case.batch_size % 2 == 0, "microbatch overlap requires even batch_size" + + init_distributed_env(model_kvargs) + dist_group_manager.create_groups(group_size=group_size) + model_cfg, _ = PretrainedConfig.get_config_dict(args.model_dir) + dist.barrier() + + torch.cuda.empty_cache() + model, _ = get_model(model_cfg, model_kvargs) + cases = resolve_batch_max_prefill_cases(args, cases, model.mem_manager.size) + cases = resolve_profile_decode_cases(args, cases, model.mem_manager.size) + cases = filter_capacity_decode_cases(args, cases, model.mem_manager.size) + if not cases: + raise ValueError("no benchmark cases remain after capacity filtering") + draft_models = init_mtp_draft_models(args, model_kvargs, model) + token_source = TokenSource(args) + executor = StaticBenchmarkExecutor(args, model, draft_models, token_source) + + results = [] + log_progress = rank_id == args.node_rank * args.tp + for case_index, case in enumerate(cases, start=1): + if log_progress: + print(f"[rank {rank_id}] case {case_index}/{len(cases)} start {case.name}", flush=True) + if args.warmup_iters > 0: + executor.run_case(case, warmup=True) + result = executor.run_case(case, warmup=False) + results.append(asdict(result)) + if log_progress: + itl = "" if result.inter_token_latency_ms is None else f" itl_ms={result.inter_token_latency_ms:.3f}" + print( + f"[rank {rank_id}] case {case_index}/{len(cases)} done elapsed_ms={result.elapsed_ms:.3f}{itl}", + flush=True, + ) + dist.barrier() + + ans_queue.put({"ok": True, "rank": rank_id, "results": results}) + except Exception: + ans_queue.put({"ok": False, "rank": rank_id, "traceback": traceback.format_exc()}) + finally: + try: + ans_queue.close() + ans_queue.join_thread() + except Exception: + pass + os._exit(0) + + +def fmt_optional(value, precision: int = 2) -> str: + if value is None: + return "-" + if isinstance(value, float): + return f"{value:.{precision}f}" + return str(value) + + +def print_aligned_table(headers: Sequence[str], rows: Sequence[Sequence[str]]): + """Print a compact right-aligned ASCII table.""" + if not rows: + return + widths = [len(str(header)) for header in headers] + for row in rows: + for index, value in enumerate(row): + widths[index] = max(widths[index], len(str(value))) + + def format_row(row: Sequence[str]) -> str: + return " ".join(str(value).rjust(widths[index]) for index, value in enumerate(row)) + + print(format_row(headers), flush=True) + print(" ".join("-" * width for width in widths), flush=True) + for row in rows: + print(format_row(row), flush=True) + + +def prefill_table_row(result: BenchmarkResult) -> List[str]: + """Format one prefill result row for stdout table output.""" + return [ + str(result.context_len), + f"{result.cache_hit_rate:.2f}", + str(result.batch_size), + fmt_optional(result.profiled_max_total_token_num, 0), + fmt_optional(result.prefill_uncached_len, 0), + fmt_optional(result.prefill_cached_len, 0), + str(result.measured_tokens), + f"{result.elapsed_ms:.3f}", + f"{result.qps:.2f}", + f"{result.tps:.2f}", + fmt_optional(result.logical_tps, 2), + ] + + +def decode_table_row(result: BenchmarkResult) -> List[str]: + """Format one decode result row for stdout table output.""" + return [ + str(result.context_len), + str(result.batch_size), + fmt_optional(result.mtp_accept_rate, 2), + fmt_optional(result.profiled_max_total_token_num, 0), + f"{result.elapsed_ms:.3f}", + f"{result.qps:.2f}", + f"{result.tps:.2f}", + fmt_optional(result.inter_token_latency_ms, 3), + ] + + +def print_results_table(results: Sequence[BenchmarkResult]): + """Print separate prefill/decode tables for measured results.""" + prefill_rows = [prefill_table_row(result) for result in results if result.stage == "prefill"] + decode_rows = [decode_table_row(result) for result in results if result.stage == "decode"] + + if prefill_rows: + print("\n[prefill]", flush=True) + print_aligned_table(PREFILL_TABLE_HEADERS, prefill_rows) + if decode_rows: + print("\n[decode]", flush=True) + print_aligned_table(DECODE_TABLE_HEADERS, decode_rows) + + +def _dp_size(args: SimpleNamespace) -> int: + return max(1, int(args.dp or 1)) + + +def _dp_world_size(args: SimpleNamespace) -> int: + return max(1, int(args.tp) // _dp_size(args)) + + +def _is_dp_group_leader(args: SimpleNamespace, rank_id: int) -> bool: + return rank_id % _dp_world_size(args) == 0 + + +def _raw_decode_step_count(result: Dict) -> int: + inter_token_latency_ms = result.get("inter_token_latency_ms") + if inter_token_latency_ms is None or inter_token_latency_ms <= 0: + return 0 + return max(1, int(round(float(result["elapsed_ms"]) / float(inter_token_latency_ms)))) + + +def aggregate_rank_results(args: SimpleNamespace, messages: Sequence[Dict]) -> List[Dict]: + """Aggregate rank-local measurements into one global result per case.""" + by_case: Dict[int, List[Dict]] = {} + for message in messages: + rank_id = int(message["rank"]) + for case_index, result in enumerate(message.get("results") or []): + by_case.setdefault(case_index, []).append({"rank": rank_id, "result": result}) + + aggregated_results: List[Dict] = [] + iters = int(args.bench_iters) + for case_index in sorted(by_case): + rank_items = sorted(by_case[case_index], key=lambda item: int(item["rank"])) + leader_items = [item for item in rank_items if _is_dp_group_leader(args, int(item["rank"]))] + if not leader_items: + leader_items = rank_items + + first = dict(leader_items[0]["result"]) + elapsed_ms = max(float(item["result"]["elapsed_ms"]) for item in rank_items) + elapsed_s = elapsed_ms / 1000.0 + batch_size = sum(int(item["result"]["batch_size"]) for item in leader_items) + measured_tokens = sum(int(item["result"]["measured_tokens"]) for item in leader_items) + + first["batch_size"] = batch_size + first["elapsed_ms"] = elapsed_ms + first["measured_tokens"] = measured_tokens + first["qps"] = batch_size * iters / elapsed_s if elapsed_s > 0 else 0.0 + first["tps"] = measured_tokens / elapsed_s if elapsed_s > 0 else 0.0 + + profiled_values = [ + int(item["result"]["profiled_max_total_token_num"]) + for item in leader_items + if item["result"].get("profiled_max_total_token_num") is not None + ] + if profiled_values: + first["profiled_max_total_token_num"] = sum(profiled_values) + + if first["stage"] == "prefill": + logical_tokens = batch_size * int(first["context_len"]) * iters + first["logical_tps"] = logical_tokens / elapsed_s if elapsed_s > 0 else 0.0 + + slowest_item = max(rank_items, key=lambda item: float(item["result"]["elapsed_ms"])) + decode_step_count = _raw_decode_step_count(slowest_item["result"]) + if decode_step_count > 0: + first["inter_token_latency_ms"] = elapsed_ms / decode_step_count + + ttft_values = [ + float(item["result"]["ttft_ms"]) for item in rank_items if item["result"].get("ttft_ms") is not None + ] + if ttft_values: + first["ttft_ms"] = max(ttft_values) + + aggregated_results.append(first) + + return aggregated_results + + +def run_benchmark(args: SimpleNamespace, cases: Sequence[BenchmarkCase]) -> List[Dict]: + ctx = mp.get_context("spawn") + ans_queue = ctx.Queue() + workers = [] + rank_start = args.node_rank * args.tp + rank_end = (args.node_rank + 1) * args.tp + case_dicts = [asdict(case) for case in cases] + args_dict = vars(args) + + for rank_id in range(rank_start, rank_end): + proc = ctx.Process(target=run_worker, args=(args_dict, case_dicts, rank_id, ans_queue)) + proc.start() + workers.append(proc) + + messages = [] + while len(messages) < len(workers): + try: + messages.append(ans_queue.get(timeout=5)) + continue + except queue.Empty: + if all(not proc.is_alive() for proc in workers): + break + + for proc in workers: + proc.join() + + failed = [message for message in messages if not message.get("ok")] + reported_ranks = {int(message["rank"]) for message in messages if "rank" in message} + failed.extend( + { + "ok": False, + "rank": rank_start + index, + "traceback": f"worker exited with code {proc.exitcode}", + } + for index, proc in enumerate(workers) + if proc.exitcode not in (0, None) + ) + failed.extend( + { + "ok": False, + "rank": rank_start + index, + "traceback": "worker did not report a result", + } + for index, proc in enumerate(workers) + if rank_start + index not in reported_ranks and proc.exitcode in (0, None) + ) + if failed: + for item in failed: + print( + f"rank {item.get('rank')} failed:\n{item.get('traceback')}", + file=sys.stderr, + ) + raise RuntimeError(f"{len(failed)} worker(s) failed") + + results = aggregate_rank_results(args, messages) + result_objs = [BenchmarkResult(**result) for result in results] + print_results_table(result_objs) + return results + + +def add_static_benchmark_args(parser: argparse.ArgumentParser): + parser.add_argument("--benchmark", choices=["all", "prefill", "decode"], default="all") + parser.add_argument("--batch_size", type=int, default=None, help="legacy single batch size") + parser.add_argument( + "--batch_sizes", + type=str, + default=None, + help="comma/space separated batch sizes", + ) + parser.add_argument("--input_len", type=int, default=64, help="legacy single prefill/context length") + parser.add_argument( + "--input_lens", + type=str, + default=None, + help="comma/space separated prefill lengths", + ) + parser.add_argument( + "--context_lens", + type=str, + default=None, + help="comma/space separated decode context lengths", + ) + parser.add_argument("--output_len", type=int, default=512, help="legacy single decode output length") + parser.add_argument( + "--output_lens", + type=str, + default=None, + help="comma/space separated decode output lengths", + ) + parser.add_argument( + "--chunked_prefill_sizes", + type=str, + default=4096, + help=("comma/space separated prefill chunk sizes; default is 4096 " "(full/none/0 select unchunked prefill)"), + ) + parser.add_argument( + "--prefill_cache_hit_rates", + type=str, + default=None, + help=( + "comma/space separated cache hit rates for prefill, e.g. " + "'0,0.5,0.8,0.9'; uncached tokens are ceil(input_len * (1-hit))" + ), + ) + parser.add_argument( + "--max_batch_size", + type=int, + default=2048, + help="upper bound for auto-computed prefill/decode batch size; <=0 disables it", + ) + parser.add_argument( + "--decode_batch_size_mode", + choices=["explicit", "profile"], + default="explicit", + help=( + "explicit uses --batch_size/--batch_sizes; profile computes decode BS " + "from profiled max_total_token_num per context" + ), + ) + parser.add_argument( + "--decode_filter_capacity", + action="store_true", + help="drop explicit decode cases whose batch size cannot fit profiled KV capacity", + ) + parser.add_argument( + "--mtp_accept_rate", + type=float, + default=1.0, + help=("per-draft-token MTP acceptance probability; sampling is outside " "the timed decode section"), + ) + parser.add_argument("--warmup_iters", type=int, default=1) + parser.add_argument("--bench_iters", type=int, default=1) + parser.add_argument("--seed", type=int, default=1234) + parser.add_argument("--dump_file", type=str, default=None, help="write aggregated benchmark results as JSON") + + +def main(argv: Optional[Sequence[str]] = None): + parser = make_argument_parser() + add_static_benchmark_args(parser) + args = parser.parse_args(argv) + if args.benchmark in {"all", "prefill"} and args.batch_max_tokens is None: + args.batch_max_tokens = 8192 + cases = build_cases(args) + if not cases: + raise ValueError("no benchmark cases were generated") + args = normalize_args(args, cases) + set_env_start_args(args) + + results = run_benchmark(args, cases) + if args.dump_file: + dump_path = Path(args.dump_file) + dump_path.parent.mkdir(parents=True, exist_ok=True) + payload = {"args": vars(args), "results": results} + dump_path.write_text(json.dumps(payload, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + mp.set_start_method("spawn", force=True) + main() From 0a48e27bcbfa8fec91f52bebb874616e1a2a95a0 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 15 Jul 2026 03:48:20 +0000 Subject: [PATCH 082/214] static benchmark support multi nodes --- .../static_inference/static_benchmark.py | 141 +++++++++++------- 1 file changed, 85 insertions(+), 56 deletions(-) diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index 3fd3c62a3e..c9bdb24d59 100644 --- a/test/benchmark/static_inference/static_benchmark.py +++ b/test/benchmark/static_inference/static_benchmark.py @@ -28,6 +28,8 @@ sys.path.append(str(REPO_ROOT)) from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4MemoryManager +from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.models import get_model from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel from lightllm.models.deepseek_v4_mtp.model import DeepseekV4MTPModel @@ -437,49 +439,38 @@ def _materialize_cached_prefix(self, req_idx: torch.Tensor, cached_len: int): req_idx_gpu = req_idx.cuda(non_blocking=True) mem_indexes_gpu = mem_indexes.reshape(batch_size, cached_len).cuda(non_blocking=True) self.model.req_manager.req_to_token_indexs[req_idx_gpu, :cached_len] = mem_indexes_gpu - self._materialize_cached_prefix_extra_slots(req_idx, cached_len) + self._materialize_cached_prefix_extra_slots(req_idx, cached_len, mem_indexes_gpu) - def _materialize_cached_prefix_extra_slots(self, req_idx: torch.Tensor, cached_len: int): + def _materialize_cached_prefix_extra_slots( + self, req_idx: torch.Tensor, cached_len: int, mem_indexes_gpu: torch.Tensor + ): req_manager = self.model.req_manager + if not isinstance(req_manager, DeepseekV4ReqManager): + return batch_size = int(req_idx.shape[0]) - b_req_idx = req_idx.cuda(non_blocking=True) - b_seq_len_cpu = cpu_i32_full((batch_size,), cached_len) - b_seq_len = b_seq_len_cpu.cuda(non_blocking=True) - - if hasattr(req_manager, "prepare_prefill_compress_slots"): - b_ready_cache_len_cpu = cpu_i32_zeros(batch_size) - req_manager.prepare_prefill_compress_slots( - b_req_idx=b_req_idx, - b_ready_cache_len=b_ready_cache_len_cpu.cuda(non_blocking=True), - b_seq_len=b_seq_len, - b_req_idx_cpu=req_idx, - b_ready_cache_len_cpu=b_ready_cache_len_cpu, - b_seq_len_cpu=b_seq_len_cpu, - ) - - if hasattr(req_manager, "prepare_prefill_swa"): - swa_ready_len = self._cached_prefix_swa_ready_len(cached_len) - b_ready_cache_len_cpu = cpu_i32_full((batch_size,), swa_ready_len) - req_manager.prepare_prefill_swa( - b_req_idx=b_req_idx, - b_ready_cache_len=b_ready_cache_len_cpu.cuda(non_blocking=True), - b_seq_len=b_seq_len, - b_req_idx_cpu=req_idx, - b_ready_cache_len_cpu=b_ready_cache_len_cpu, - b_seq_len_cpu=b_seq_len_cpu, - ) + req_list = req_idx.tolist() + seq_list = [cached_len] * batch_size + + swa_ready_len = self._cached_prefix_swa_ready_len(cached_len) + req_manager.prepare_prefill_swa( + req_list=req_list, + ready_list=[swa_ready_len] * batch_size, + seq_list=seq_list, + mem_indexes=mem_indexes_gpu[:, swa_ready_len:].contiguous(), + ) + req_manager.prepare_prefill_compress_slots( + req_list=req_list, + ready_list=[0] * batch_size, + seq_list=seq_list, + mem_indexes=mem_indexes_gpu, + ) def _cached_prefix_swa_ready_len(self, cached_len: int) -> int: - req_manager = self.model.req_manager - retain_fn = getattr(req_manager, "_swa_retain_len", None) - if retain_fn is not None: - retain_len = int(retain_fn()) - else: - retain_len = int(getattr(req_manager, "sliding_window", cached_len) or cached_len) - ready_len = max(0, int(cached_len) - max(1, retain_len)) - page_fn = getattr(req_manager, "get_prompt_cache_page_size", None) - page_size = int(page_fn()) if page_fn is not None else 1 - return ready_len // max(1, page_size) * max(1, page_size) + req_manager: DeepseekV4ReqManager = self.model.req_manager + retain_len = int(req_manager._swa_retain_len()) + ready_len = max(0, int(cached_len) - retain_len) + page_size = int(req_manager.get_prompt_cache_page_size()) + return ready_len // page_size * page_size def _make_prefill_input(self, token_chunk: np.ndarray, req_idx: torch.Tensor, ready_cache_len: int) -> ModelInput: batch_size, q_len = token_chunk.shape @@ -785,11 +776,11 @@ def apply_max_batch_size(batch_size: int, max_batch_size: int) -> int: def prefill_batch_size_from_batch_max_tokens( batch_max_tokens: int, - step_tokens_per_req: int, + uncached_tokens_per_req: int, max_batch_size: int, ) -> int: - """Compute prefill BS from batch_max_tokens before KV-capacity capping.""" - batch_size = max(1, int(batch_max_tokens) // max(1, step_tokens_per_req)) + """Compute prefill BS from the full uncached suffix before KV-capacity capping.""" + batch_size = max(1, int(batch_max_tokens) // max(1, uncached_tokens_per_req)) return apply_max_batch_size(batch_size, max_batch_size) @@ -799,7 +790,7 @@ def build_prefill_cases( chunk_sizes: Sequence[Optional[int]], cache_hit_rates: Sequence[float], ) -> List[BenchmarkCase]: - """Build full-prefill cases using batch_max_tokens per chunk step.""" + """Build full-prefill cases using batch_max_tokens per uncached suffix.""" if args.batch_max_tokens is None: raise ValueError("prefill benchmark requires --batch_max_tokens") cases: List[BenchmarkCase] = [] @@ -810,7 +801,7 @@ def build_prefill_cases( step_tokens = prefill_step_tokens_per_req(uncached_len, chunk_size) bs = prefill_batch_size_from_batch_max_tokens( args.batch_max_tokens, - step_tokens, + uncached_len, args.max_batch_size, ) chunk_name = chunk_size if chunk_size else "none" @@ -867,7 +858,8 @@ def build_cases(args: SimpleNamespace) -> List[BenchmarkCase]: input_lens = parse_int_list(args.input_lens, [args.input_len]) context_lens = parse_int_list(args.context_lens, input_lens) output_lens = parse_int_list(args.output_lens, [args.output_len]) - chunk_sizes = parse_chunk_sizes(args.chunked_prefill_sizes, args.chunked_prefill_size) + fallback_chunk_size = args.chunked_prefill_size if args.chunked_prefill_size is not None else 4096 + chunk_sizes = parse_chunk_sizes(args.chunked_prefill_sizes, fallback_chunk_size) cache_hit_rates = parse_float_list(args.prefill_cache_hit_rates, [0.0]) cases: List[BenchmarkCase] = [] @@ -887,19 +879,25 @@ def decode_profile_batch_divisor(args: SimpleNamespace, case: BenchmarkCase) -> def filter_capacity_decode_cases( args: SimpleNamespace, cases: Sequence[BenchmarkCase], - profiled_max_total_token_num: int, + mem_manager, ) -> List[BenchmarkCase]: if not getattr(args, "decode_filter_capacity", False): return list(cases) resolved: List[BenchmarkCase] = [] - capacity_tokens = int(profiled_max_total_token_num) + capacity_tokens = int(mem_manager.size) for case in cases: if case.stage != "decode": resolved.append(case) continue divisor = decode_profile_batch_divisor(args, case) - if case.batch_size * divisor <= capacity_tokens: + fits_capacity = case.batch_size * divisor <= capacity_tokens + if fits_capacity and isinstance(mem_manager, DeepseekV4MemoryManager) and mem_manager.n_c4 > 0: + c4_page_size = int(mem_manager.c4_pool.page_size) + c4_entries_per_req = divisor // 4 + c4_pages_per_req = (c4_entries_per_req + c4_page_size - 1) // c4_page_size + fits_capacity = case.batch_size * c4_pages_per_req <= mem_manager.c4_num_pages + if fits_capacity: resolved.append( replace( case, @@ -993,6 +991,13 @@ def resolve_batch_max_prefill_cases( def normalize_args(args: argparse.Namespace, cases: Sequence[BenchmarkCase]) -> SimpleNamespace: """Fill LightLLM startup args needed before model construction.""" + if args.nnodes <= 0: + raise ValueError(f"--nnodes must be positive, got {args.nnodes}") + if not 0 <= args.node_rank < args.nnodes: + raise ValueError(f"--node_rank must be in [0, {args.nnodes}), got {args.node_rank}") + if args.tp % args.nnodes != 0: + raise ValueError(f"--tp must be divisible by --nnodes, got tp={args.tp} nnodes={args.nnodes}") + if args.data_type is None: args.data_type = get_dtype(args.model_dir) @@ -1034,7 +1039,12 @@ def normalize_args(args: argparse.Namespace, cases: Sequence[BenchmarkCase]) -> and args.max_total_token_num is None ) prefill_batch_size_needs_profile = args.benchmark in {"all", "prefill"} and args.max_total_token_num is None - needs_profiled_batch_size = decode_batch_size_needs_profile or prefill_batch_size_needs_profile + decode_capacity_needs_profile = ( + args.benchmark in {"all", "decode"} and args.decode_filter_capacity and args.max_total_token_num is None + ) + needs_profiled_batch_size = ( + decode_batch_size_needs_profile or prefill_batch_size_needs_profile or decode_capacity_needs_profile + ) if args.max_total_token_num is None and not needs_profiled_batch_size: args.max_total_token_num = max_batch * (args.max_req_total_len + mtp_width + 8) @@ -1200,7 +1210,7 @@ def run_worker(args_dict: Dict, case_dicts: List[Dict], rank_id: int, ans_queue) model, _ = get_model(model_cfg, model_kvargs) cases = resolve_batch_max_prefill_cases(args, cases, model.mem_manager.size) cases = resolve_profile_decode_cases(args, cases, model.mem_manager.size) - cases = filter_capacity_decode_cases(args, cases, model.mem_manager.size) + cases = filter_capacity_decode_cases(args, cases, model.mem_manager) if not cases: raise ValueError("no benchmark cases remain after capacity filtering") draft_models = init_mtp_draft_models(args, model_kvargs, model) @@ -1208,7 +1218,8 @@ def run_worker(args_dict: Dict, case_dicts: List[Dict], rank_id: int, ans_queue) executor = StaticBenchmarkExecutor(args, model, draft_models, token_source) results = [] - log_progress = rank_id == args.node_rank * args.tp + local_world_size = args.tp // args.nnodes + log_progress = rank_id == args.node_rank * local_world_size for case_index, case in enumerate(cases, start=1): if log_progress: print(f"[rank {rank_id}] case {case_index}/{len(cases)} start {case.name}", flush=True) @@ -1224,7 +1235,13 @@ def run_worker(args_dict: Dict, case_dicts: List[Dict], rank_id: int, ans_queue) ) dist.barrier() - ans_queue.put({"ok": True, "rank": rank_id, "results": results}) + message = {"ok": True, "rank": rank_id, "results": results} + if args.nnodes > 1: + global_messages = [None] * args.tp + dist.all_gather_object(global_messages, message) + if rank_id == 0: + message["global_messages"] = global_messages + ans_queue.put(message) except Exception: ans_queue.put({"ok": False, "rank": rank_id, "traceback": traceback.format_exc()}) finally: @@ -1385,8 +1402,9 @@ def run_benchmark(args: SimpleNamespace, cases: Sequence[BenchmarkCase]) -> List ctx = mp.get_context("spawn") ans_queue = ctx.Queue() workers = [] - rank_start = args.node_rank * args.tp - rank_end = (args.node_rank + 1) * args.tp + local_world_size = args.tp // args.nnodes + rank_start = args.node_rank * local_world_size + rank_end = rank_start + local_world_size case_dicts = [asdict(case) for case in cases] args_dict = vars(args) @@ -1435,6 +1453,14 @@ def run_benchmark(args: SimpleNamespace, cases: Sequence[BenchmarkCase]) -> List ) raise RuntimeError(f"{len(failed)} worker(s) failed") + if args.nnodes > 1: + if args.node_rank != 0: + return [] + gathered = [message.get("global_messages") for message in messages if message.get("global_messages")] + if len(gathered) != 1 or len(gathered[0]) != args.tp: + raise RuntimeError("rank 0 did not receive one result payload from every global rank") + messages = gathered[0] + results = aggregate_rank_results(args, messages) result_objs = [BenchmarkResult(**result) for result in results] print_results_table(result_objs) @@ -1473,8 +1499,11 @@ def add_static_benchmark_args(parser: argparse.ArgumentParser): parser.add_argument( "--chunked_prefill_sizes", type=str, - default=4096, - help=("comma/space separated prefill chunk sizes; default is 4096 " "(full/none/0 select unchunked prefill)"), + default=None, + help=( + "comma/space separated prefill chunk sizes; overrides --chunked_prefill_size; " + "both omitted defaults to 4096 (full/none/0 select unchunked prefill)" + ), ) parser.add_argument( "--prefill_cache_hit_rates", @@ -1530,7 +1559,7 @@ def main(argv: Optional[Sequence[str]] = None): set_env_start_args(args) results = run_benchmark(args, cases) - if args.dump_file: + if args.dump_file and args.node_rank == 0: dump_path = Path(args.dump_file) dump_path.parent.mkdir(parents=True, exist_ok=True) payload = {"args": vars(args), "results": results} From ad4508aaa6c5092256e7fddfff1afe8942016961 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 17 Jul 2026 05:51:32 +0000 Subject: [PATCH 083/214] make c128 compression state request-local and MTP-safe --- .../deepseek4_mem_manager.py | 30 ++++++----- .../kv_cache_mem_manager/mem_manager.py | 13 ++++- lightllm/common/req_manager.py | 12 ++--- .../deepseek_v4/layer_infer/compressor.py | 50 ++++++++++--------- lightllm/models/deepseek_v4/model.py | 1 + .../server/router/model_infer/infer_batch.py | 11 ++-- 6 files changed, 69 insertions(+), 48 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index d07406529b..c39cd8ef1f 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -30,9 +30,10 @@ DSV4_C4_PAGE_SIZE = 64 # 64 slots/page DSV4_C128_PAGE_SIZE = 2 # 2 slots/page DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 # 256 (= c4 ratio) -# compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 离线聚合为每页 1 组。 +# compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 是每请求的基础窗口, +# 实际宽度会追加 mtp_step 个候选槽,再向上对齐到 ratio 4。 DSV4_C4_STATE_RING = 8 # 8 rows/page -DSV4_C128_STATE_RING = 128 # 128 rows/page +DSV4_C128_STATE_RING = 128 # 128 rows/request before MTP padding # swa 池占 full token 空间的比例(sglang DSV4 默认 swa_full_tokens_ratio=0.1 同值)。 # 瞬时借页/驱逐走 swa 压力阀;池子大小仅按 ratio 切分,不再叠加结构性余量。 DSV4_SWA_FULL_TOKENS_RATIO = 0.1 # 0.1 @@ -148,8 +149,9 @@ def __init__( head_dim, layer_num, compress_rates: List[int], + max_request_num: int, + mtp_step: int, indexer_head_dim: int = 128, - max_request_num: Optional[int] = None, swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, @@ -161,12 +163,15 @@ def __init__( ), f"DeepSeek-V4 packed indexer-K 期望 indexer_head_dim={self.indexer_head_dim_default}" assert len(compress_rates) == layer_num, f"compress_rates 长度 {len(compress_rates)} 必须等于 layer_num {layer_num}" assert all(r in (0, 4, 128) for r in compress_rates), "compress_rates 取值只能是 0/4/128" + assert max_request_num > 0, "max_request_num 必须为正数" + assert 0 <= mtp_step < DSV4_C128_STATE_RING, "mtp_step 必须位于 [0, 128)" self.compress_rates = list(compress_rates) self.n_c4 = sum(1 for r in self.compress_rates if r == 4) self.n_c128 = sum(1 for r in self.compress_rates if r == 128) self.indexer_head_dim = indexer_head_dim self.max_request_num = max_request_num + self.c128_state_ring = _ceil_div(DSV4_C128_STATE_RING + mtp_step, 4) * 4 self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) @@ -204,16 +209,16 @@ def get_cell_size(self): indexer_bytes = self.indexer_bytes_per_token state_dtype_bytes = torch._utils._element_size(torch.float32) c4_state_width = 4 * self.head_dim + 4 * self.indexer_head_dim - c128_state_width = 2 * self.head_dim c4_state_bytes = DSV4_C4_STATE_RING / DSV4_SWA_PAGE_SIZE * c4_state_width * state_dtype_bytes * self.n_c4 - c128_state_bytes = ( - DSV4_C128_STATE_RING / DSV4_SWA_PAGE_SIZE * c128_state_width * state_dtype_bytes * self.n_c128 - ) - swa_slot = kv_bytes * self.layer_num + c4_state_bytes + c128_state_bytes + swa_slot = kv_bytes * self.layer_num + c4_state_bytes compressed = (kv_bytes + indexer_bytes) * self.n_c4 / 4 + kv_bytes * self.n_c128 / 128 return swa_slot * self.swa_full_tokens_ratio + compressed + def get_fixed_memory_size(self): + state_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 + return self.n_c128 * state_rows * (2 * self.head_dim) * torch._utils._element_size(torch.float32) + # ------------------------------------------------------------------ buffers def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): rank_in_node = get_current_rank_in_node() @@ -312,9 +317,9 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): ) self.full_to_c128_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") self.full_to_c128_indexs[size] = self.c128_pool.HOLD_TOKEN_MEMINDEX - # c128 compressor 在途状态: 与 c4 同样由 full->swa 推导行号,但 ring=128 且无 overlap。 - # last_dim = 2*head_dim;末行是 swa 缺失/出窗时读取的哨兵。 - state_rows = self._paged_state_rows(self.swa_num_pages, DSV4_C128_STATE_RING, 128) + # c128 compressor 在途状态按 request 寻址。每个 request 保留完整的 128-token + # 聚合窗口以及 MTP 候选余量;最后一行供无效请求/位置读取哨兵。 + state_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 self.c128_state_buffer = torch.zeros( (self.n_c128, state_rows, 2 * self.head_dim), dtype=torch.float32, device="cuda" ) @@ -323,6 +328,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): logger.info( f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " + f"c128_state_ring={self.c128_state_ring} max_requests={self.max_request_num} " f"packed_kv_bytes={self.mla_bytes_per_token} indexer_bytes={self.indexer_bytes_per_token}" ) @@ -355,7 +361,7 @@ def get_c4_indexer_state_buffer(self, layer_index: int) -> torch.Tensor: return self.c4_indexer_state_buffer[self.layer_to_c4_idx[layer_index]] def get_c128_state_buffer(self, layer_index: int) -> torch.Tensor: - assert self.compress_rates[layer_index] == 128, "只有 c128(HCA) 层有 paged compressor state" + assert self.compress_rates[layer_index] == 128, "只有 c128(HCA) 层有 request-scoped compressor state" return self.c128_state_buffer[self.layer_to_c128_idx[layer_index]] # ------------------------------------------------------------------ swa slot lifecycle diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 658d3e899c..4aa0937335 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -58,6 +58,9 @@ def get_att_input_params(self, layer_index: int) -> Tuple[Any, Any]: def get_cell_size(self): return 2 * self.head_num * self.head_dim * self.layer_num * torch._utils._element_size(self.dtype) + def get_fixed_memory_size(self): + return 0 + def profile_size(self, mem_fraction): if self.size is not None: return @@ -66,13 +69,21 @@ def profile_size(self, mem_fraction): world_size = dist.get_world_size() available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) cell_size = self.get_cell_size() - self.size = int(available_memory * 1024 ** 3 / cell_size) + fixed_memory_size = self.get_fixed_memory_size() + available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size + if available_memory_bytes <= 0: + raise RuntimeError( + f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " + f"but only {available_memory:.2f} GB is available for KV cache" + ) + self.size = int(available_memory_bytes / cell_size) if world_size > 1: tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}") dist.all_reduce(tensor, op=dist.ReduceOp.MIN) self.size = tensor.item() logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" + f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 8eb44b60d7..64b5e361ab 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -33,9 +33,9 @@ class DeepseekV4PromptCachePayload: """prompt cache 载荷: swa 按页有效性 bitmap 和最后有效页。 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 - 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4/c128 compressor 状态以 swa - 页派生寻址(随 swa 页生灭,命中零拷贝续算)。prompt cache 对齐到 256 token, - 避免共享前缀停在 c4 物理页中间。 + 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4 compressor 状态随 swa 页 + 生灭。c128 状态按 request ring 寻址;prompt cache 的 256-token 边界同时是 c128 + 分组边界,命中后新分组会在首次读取前覆写完整 128-token 窗口,因而无需保存状态。 * ``swa_page_valid``: cpu bool [cache_len // page],插入时按当下 full_to_swa 映射写定 (页内 token 映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 page 边界, @@ -562,7 +562,7 @@ def prepare_decode_swa( def init_compress_state(self, req_idx: int): """新请求开始时重置 runtime 水位线(对应 mamba 的 init_linear_att_state 调用点)。 - c4/c128 compressor state 都随 swa 页寻址,由内核按组覆写;请求复用时不做 per-req 清零。""" + c4 状态随 swa 页寻址;c128 request ring 依靠 overwrite-before-read,不做大块清零。""" self.clear_runtime_state(req_idx) return @@ -875,8 +875,8 @@ def build_prompt_cache_payload( self, cache_len: int, ) -> DeepseekV4PromptCachePayload: - """构造插入载荷。compressor 状态不进载荷(c4 随 swa 页生灭、c128 边界自然归零), - cache_len 不再受序列末端约束——任意 128 对齐前缀皆可插入。 + """构造插入载荷。compressor 状态不进载荷(c4 随 swa 页生灭、c128 在 256 对齐 + 恢复点依靠 overwrite-before-read),cache_len 不再受序列末端约束。 swa_page_valid 不在此填: 它必须用插入时刻的映射(infer batch 在 insert 前补)。""" assert self.mem_manager is not None return DeepseekV4PromptCachePayload(cache_len=int(cache_len)) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index e9f2af7385..e068078004 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -7,11 +7,7 @@ from triton.language.extra import libdevice from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( - DSV4_C4_STATE_RING, - DSV4_C128_STATE_RING, - DSV4_SWA_PAGE_SIZE, -) +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C4_STATE_RING, DSV4_SWA_PAGE_SIZE @dataclass @@ -21,10 +17,12 @@ class CoreCompressorMetadata: out_slots: torch.Tensor mem_index: torch.Tensor state_buffer: torch.Tensor + state_ring: int out_buffer: torch.Tensor out_page_size: int position_ids: torch.Tensor b_req_idx: torch.Tensor + b_mtp_index: torch.Tensor b_seq_len: torch.Tensor b_ready_cache_len: Optional[torch.Tensor] b_q_start_loc: Optional[torch.Tensor] @@ -97,11 +95,15 @@ def _save_partial_states_kernel( if position + COMPRESS_RATIO < seq_len: return - full_slot = tl.load(mem_index + token_idx).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) - if swa_slot < 0: - return - state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) + if IS_C4: + full_slot = tl.load(mem_index + token_idx).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + if swa_slot < 0: + return + state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) + else: + req_idx = tl.load(b_req_idx + batch_idx).to(tl.int64) + state_row = req_idx * STATE_RING + position % STATE_RING offs = tl.arange(0, BLOCK) mask = offs < STATE_WIDTH @@ -122,6 +124,7 @@ def _fused_compress_norm_rope_insert_kernel( positions, token_to_batch_idx, b_req_idx, + b_mtp_index, b_seq_len, b_ready_cache_len, b_q_start_loc, @@ -174,15 +177,16 @@ def _fused_compress_norm_rope_insert_kernel( ready_len = tl.load(b_ready_cache_len + batch_idx) q_start = tl.load(b_q_start_loc + batch_idx) else: - ready_len = position - q_start = token_idx + mtp_index = tl.load(b_mtp_index + batch_idx) + ready_len = position - mtp_index + q_start = token_idx - mtp_index token_offsets = tl.arange(0, WINDOW_SIZE) start = position - WINDOW_SIZE + 1 gather_pos = start + token_offsets valid_pos = (gather_pos >= 0) & (gather_pos < seq_len) - use_current = (gather_pos >= ready_len) & valid_pos if IS_PREFILL else gather_pos == position - current_idx = q_start + (gather_pos - ready_len) if IS_PREFILL else token_idx + token_offsets * 0 + use_current = (gather_pos >= ready_len) & valid_pos + current_idx = q_start + (gather_pos - ready_len) if IS_C4: full_slot = tl.load( @@ -195,14 +199,8 @@ def _fused_compress_norm_rope_insert_kernel( state_valid = valid_pos & (~use_current) & (swa_slot >= 0) head_offset = tl.where(token_offsets >= COMPRESS_RATIO, HEAD_DIM, 0) else: - full_slot = tl.load( - req_to_token + req_idx * req_to_token_stride0 + gather_pos, - mask=valid_pos & (~use_current), - other=0, - ).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot, mask=valid_pos & (~use_current), other=-1).to(tl.int64) - state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) - state_valid = valid_pos & (~use_current) & (swa_slot >= 0) + state_row = req_idx * STATE_RING + gather_pos % STATE_RING + state_valid = valid_pos & (~use_current) head_offset = token_offsets * 0 offs = tl.arange(0, BLOCK) @@ -306,6 +304,7 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) + state_ring = DSV4_C4_STATE_RING out_buffer = torch.empty( (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), dtype=torch.bfloat16, @@ -316,10 +315,12 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, if compress_ratio == 4: out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] state_buffer = mem_manager.get_c4_state_buffer(layer_idx) + state_ring = DSV4_C4_STATE_RING out_pool = mem_manager.c4_pool elif compress_ratio == 128: out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] state_buffer = mem_manager.get_c128_state_buffer(layer_idx) + state_ring = mem_manager.c128_state_ring out_pool = mem_manager.c128_pool else: raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") @@ -343,10 +344,12 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, out_slots=out_slots, mem_index=infer_state.mem_index, state_buffer=state_buffer, + state_ring=state_ring, out_buffer=out_buffer, out_page_size=out_page_size, position_ids=infer_state.position_ids, b_req_idx=infer_state.b_req_idx, + b_mtp_index=infer_state.b_mtp_index, b_seq_len=infer_state.b_seq_len, b_ready_cache_len=infer_state.b_ready_cache_len, b_q_start_loc=infer_state.b_q_start_loc, @@ -403,7 +406,7 @@ def fused_compress( state_width = kv_score.shape[-1] // 2 state_last_dim = metadata.state_buffer.shape[-1] is_c4 = compress_ratio == 4 - state_ring = DSV4_C4_STATE_RING if is_c4 else DSV4_C128_STATE_RING + state_ring = metadata.state_ring block_state = triton.next_power_of_2(state_width) block_head = triton.next_power_of_2(head_dim) @@ -415,6 +418,7 @@ def fused_compress( metadata.position_ids, metadata.token_to_batch_idx, metadata.b_req_idx, + metadata.b_mtp_index, metadata.b_seq_len, metadata.b_ready_cache_len if metadata.b_ready_cache_len is not None else metadata.b_seq_len, metadata.b_q_start_loc if metadata.b_q_start_loc is not None else metadata.b_seq_len, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index e0640a52dd..90ba83a70f 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -89,6 +89,7 @@ def _init_mem_manager(self): compress_rates=self._get_compress_rates(layer_num), indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, + mtp_step=self.args.mtp_step, mem_fraction=self.mem_fraction, ) self.req_manager.mem_manager = self.mem_manager diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 6dfc35dbc2..5e3de340cc 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -129,9 +129,8 @@ def free_a_req_mem(self, free_token_index: List, req: "InferReq"): if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.cur_kv_len]) if self.is_deepseek_v4: - # 槽位随 full 槽经 mem_manager.free 级联回收。pause 路径不释放 req_idx, - # 必须在此复位出窗水位线 + 清 c128 在途状态(恢复命中走 extend,不会再有 - # restore/zero 时机;c4 状态随 swa 页生灭,无需处理)。 + # 槽位随 full 槽经 mem_manager.free 级联回收。pause 路径不释放 req_idx, + # 这里只复位出窗水位线;恢复从 0 重算,c128 request ring 会在读取前覆写。 self.req_manager.init_compress_state(req.req_idx) else: if not self.is_linear_att_mixed_model: @@ -738,9 +737,9 @@ def _match_radix_cache(self): ready_cache_len = share_node.node_prefix_total_len # 从 cpu 到 gpu 是流内阻塞操作 g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor - # DeepSeek-V4 命中无需任何恢复: 槽位由 full_to_* 映射键控(radix 持有 full 槽即有效, - # 命中长度已在 match_prefix 内按 bitmap 裁剪),c4 compressor 状态随 swa 页常驻 - # (零拷贝续算),c128 状态在 128 对齐边界自然归零(init_compress_state 已清)。 + # DeepSeek-V4 命中无需恢复 compressor 状态: 槽位由 full_to_* 映射键控(radix + # 持有 full 槽即有效),c4 状态随 swa 页常驻;命中点按 256 token 对齐,c128 + # request ring 的下一组会在首次读取前完整覆写。 self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.shm_req.prompt_cache_len = self.cur_kv_len # 记录 prompt cache 的命中长度 From f4116714340f8eb762ba4ab84a777d08ab81ca46 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 17 Jul 2026 06:09:55 +0000 Subject: [PATCH 084/214] fix c4 mtp overwrite bug --- .../kv_cache_mem_manager/deepseek4_mem_manager.py | 14 ++++++++------ .../models/deepseek_v4/layer_infer/compressor.py | 6 +++--- 2 files changed, 11 insertions(+), 9 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index c39cd8ef1f..abc15ecd0c 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -30,9 +30,9 @@ DSV4_C4_PAGE_SIZE = 64 # 64 slots/page DSV4_C128_PAGE_SIZE = 2 # 2 slots/page DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 # 256 (= c4 ratio) -# compressor state ring: c4 overlap 对为每页 2 个分组槽 × ratio 4 行;c128 是每请求的基础窗口, -# 实际宽度会追加 mtp_step 个候选槽,再向上对齐到 ratio 4。 -DSV4_C4_STATE_RING = 8 # 8 rows/page +# compressor state ring: c4 overlap 对的基础窗口为每页 2 个分组槽 × ratio 4 行;MTP +# 追加候选槽,避免 rejected draft 覆盖仍存活的基础窗口。c128 同样追加候选槽,再对齐到 ratio 4。 +DSV4_C4_STATE_RING = 8 # 8 rows/page before MTP padding DSV4_C128_STATE_RING = 128 # 128 rows/request before MTP padding # swa 池占 full token 空间的比例(sglang DSV4 默认 swa_full_tokens_ratio=0.1 同值)。 # 瞬时借页/驱逐走 swa 压力阀;池子大小仅按 ratio 切分,不再叠加结构性余量。 @@ -171,6 +171,7 @@ def __init__( self.n_c128 = sum(1 for r in self.compress_rates if r == 128) self.indexer_head_dim = indexer_head_dim self.max_request_num = max_request_num + self.c4_state_ring = DSV4_C4_STATE_RING + mtp_step self.c128_state_ring = _ceil_div(DSV4_C128_STATE_RING + mtp_step, 4) * 4 self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) @@ -209,7 +210,7 @@ def get_cell_size(self): indexer_bytes = self.indexer_bytes_per_token state_dtype_bytes = torch._utils._element_size(torch.float32) c4_state_width = 4 * self.head_dim + 4 * self.indexer_head_dim - c4_state_bytes = DSV4_C4_STATE_RING / DSV4_SWA_PAGE_SIZE * c4_state_width * state_dtype_bytes * self.n_c4 + c4_state_bytes = self.c4_state_ring / DSV4_SWA_PAGE_SIZE * c4_state_width * state_dtype_bytes * self.n_c4 swa_slot = kv_bytes * self.layer_num + c4_state_bytes compressed = (kv_bytes + indexer_bytes) * self.n_c4 / 4 + kv_bytes * self.n_c128 / 128 @@ -294,7 +295,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): # 生灭 -> radix 命中零拷贝续算。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), # 取整到 ratio;末行哨兵 kv=0/score=-inf(KVAndScore.clear 语义),其余行由内核在 # 组起点覆写,无需按页清零。last_dim = 2*coff*head_dim(overlap coff=2)。 - state_rows = self._paged_state_rows(self.swa_num_pages, DSV4_C4_STATE_RING, 4) + state_rows = self._paged_state_rows(self.swa_num_pages, self.c4_state_ring, 4) self.c4_state_buffer = torch.zeros( (self.n_c4, state_rows, 4 * self.head_dim), dtype=torch.float32, device="cuda" ) @@ -328,7 +329,8 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): logger.info( f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " - f"c128_state_ring={self.c128_state_ring} max_requests={self.max_request_num} " + f"c4_state_ring={self.c4_state_ring} c128_state_ring={self.c128_state_ring} " + f"max_requests={self.max_request_num} " f"packed_kv_bytes={self.mla_bytes_per_token} indexer_bytes={self.indexer_bytes_per_token}" ) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index e068078004..25e1df8601 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -7,7 +7,7 @@ from triton.language.extra import libdevice from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C4_STATE_RING, DSV4_SWA_PAGE_SIZE +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE @dataclass @@ -304,7 +304,7 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) - state_ring = DSV4_C4_STATE_RING + state_ring = mem_manager.c4_state_ring out_buffer = torch.empty( (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), dtype=torch.bfloat16, @@ -315,7 +315,7 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, if compress_ratio == 4: out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] state_buffer = mem_manager.get_c4_state_buffer(layer_idx) - state_ring = DSV4_C4_STATE_RING + state_ring = mem_manager.c4_state_ring out_pool = mem_manager.c4_pool elif compress_ratio == 128: out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] From 9a071ca5e0e53bab4bca63586f78e92b493cfb68 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 17 Jul 2026 06:54:02 +0000 Subject: [PATCH 085/214] autoset NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE --- lightllm/utils/envs_utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 773320273c..bd9ce0e211 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -83,8 +83,10 @@ def get_deepep_num_max_dispatch_tokens_per_rank_prefill(): @lru_cache(maxsize=None) def get_deepep_num_max_dispatch_tokens_per_rank_decode(): - # 该参数需要大于单卡最大batch size,且是8的倍数。该参数与显存占用直接相关,值越大,显存占用越大,如果出现显存不足,可以尝试调小该值 - return int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", 256)) + args = get_env_start_args() + required = args.running_max_req_size * (args.mtp_step + 1) + required = ((required + 7) // 8) * 8 + return int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", required)) def get_lightllm_gunicorn_keep_alive(): From 18b83bca1f5ab7150d76eb1f2ba200afebf686d1 Mon Sep 17 00:00:00 2001 From: sufubao Date: Tue, 7 Jul 2026 10:00:38 +0800 Subject: [PATCH 086/214] feat: support /v1/responses --- lightllm/server/api_http.py | 21 ++ lightllm/server/api_responses.py | 558 +++++++++++++++++++++++++++++++ 2 files changed, 579 insertions(+) create mode 100644 lightllm/server/api_responses.py diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index 270e2a8cfd..b433cf748d 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -353,6 +353,27 @@ async def anthropic_messages(raw_request: Request) -> Response: return Response(status_code=499) +@app.post("/v1/responses") +async def openai_responses(raw_request: Request) -> Response: + if get_env_start_args().run_mode in ["prefill", "decode"]: + return create_error_response( + HTTPStatus.EXPECTATION_FAILED, "service in pd mode dont recv reqs from http interface" + ) + from .api_responses import responses_impl + + try: + return await responses_impl(raw_request) + except ServerBusyError as e: + logger.error("%s", str(e), exc_info=True) + return create_error_response(HTTPStatus.SERVICE_UNAVAILABLE, str(e)) + except ClientDisconnected as e: + logger.warning(str(e)) + return Response(status_code=499) + except Exception as e: + logger.error("An error occurred: %s", str(e), exc_info=True) + return create_error_response(HTTPStatus.EXPECTATION_FAILED, str(e)) + + @app.get("/v1/models", response_model=ModelListResponse) async def get_models(raw_request: Request): model_name = g_objs.args.model_name diff --git a/lightllm/server/api_responses.py b/lightllm/server/api_responses.py new file mode 100644 index 0000000000..4ce0e19673 --- /dev/null +++ b/lightllm/server/api_responses.py @@ -0,0 +1,558 @@ +from __future__ import annotations + +import time +import uuid +import ujson as json +from http import HTTPStatus +from typing import Any, AsyncGenerator, Dict, List, Optional + +from fastapi import Request +from fastapi.responses import JSONResponse, Response, StreamingResponse + +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +def _content_parts_to_chat(parts: List[Any]) -> List[Dict[str, Any]]: + chat_parts = [] + for part in parts: + if isinstance(part, str): + chat_parts.append({"type": "text", "text": part}) + continue + ptype = part.get("type") + if ptype in ("input_text", "output_text", "text"): + chat_parts.append({"type": "text", "text": part.get("text", "")}) + elif ptype == "refusal": + chat_parts.append({"type": "text", "text": part.get("refusal", "")}) + elif ptype == "input_image": + url = part.get("image_url") + if isinstance(url, dict): + url = url.get("url") + if not url: + raise ValueError("input_image requires an image_url") + chat_parts.append({"type": "image_url", "image_url": {"url": url}}) + elif ptype == "input_audio": + audio = part.get("input_audio") or {} + url = audio.get("url") or part.get("audio_url") + if not url: + raise ValueError("input_audio requires a url (raw audio data is not supported)") + chat_parts.append({"type": "audio_url", "audio_url": {"url": url}}) + else: + raise ValueError(f"Unsupported input content type: {ptype}") + return chat_parts + + +def _input_items_to_messages(items: List[Any]) -> List[Dict[str, Any]]: + messages: List[Dict[str, Any]] = [] + for item in items: + if not isinstance(item, dict): + raise ValueError("input items must be objects") + itype = item.get("type") + if itype == "function_call": + messages.append( + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": item.get("call_id") or item.get("id"), + "type": "function", + "function": {"name": item.get("name"), "arguments": item.get("arguments") or ""}, + } + ], + } + ) + elif itype == "function_call_output": + output = item.get("output") + if not isinstance(output, str): + output = json.dumps(output, ensure_ascii=False) + messages.append({"role": "tool", "tool_call_id": item.get("call_id"), "content": output}) + elif itype == "reasoning": + continue + elif itype in (None, "message"): + role = item.get("role", "user") + if role == "developer": + role = "system" + content = item.get("content") + if isinstance(content, list): + content = _content_parts_to_chat(content) + messages.append({"role": role, "content": content}) + else: + raise ValueError(f"Unsupported input item type: {itype}") + return messages + + +def _tools_to_chat(tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + chat_tools = [] + for tool in tools: + if tool.get("type") != "function": + logger.warning("Ignoring unsupported tool type: %s", tool.get("type")) + continue + if "function" in tool: + chat_tools.append(tool) + continue + chat_tools.append( + { + "type": "function", + "function": { + "name": tool.get("name"), + "description": tool.get("description"), + "parameters": tool.get("parameters"), + }, + } + ) + return chat_tools + + +def _text_format_to_response_format(body: Dict[str, Any]) -> Optional[Dict[str, Any]]: + fmt = (body.get("text") or {}).get("format") + if not fmt: + return None + ftype = fmt.get("type") + if ftype == "text": + return None + if ftype == "json_object": + return {"type": "json_object"} + if ftype == "json_schema": + return { + "type": "json_schema", + "json_schema": { + "name": fmt.get("name", "response"), + "description": fmt.get("description"), + "schema": fmt.get("schema"), + "strict": fmt.get("strict"), + }, + } + raise ValueError(f"Unsupported text.format type: {ftype}") + + +def _responses_to_chat_request(body: Dict[str, Any]) -> Dict[str, Any]: + messages: List[Dict[str, Any]] = [] + if body.get("instructions"): + messages.append({"role": "system", "content": body["instructions"]}) + + inp = body.get("input") + if isinstance(inp, str): + messages.append({"role": "user", "content": inp}) + elif isinstance(inp, list): + messages.extend(_input_items_to_messages(inp)) + else: + raise ValueError("'input' must be a string or an array of input items") + + system_texts = [] + for m in messages: + if m["role"] == "system": + content = m["content"] + if isinstance(content, list): + content = "\n".join(p.get("text", "") for p in content) + if content: + system_texts.append(content) + messages = [m for m in messages if m["role"] != "system"] + if system_texts: + messages.insert(0, {"role": "system", "content": "\n\n".join(system_texts)}) + + chat: Dict[str, Any] = { + "model": body.get("model", "default"), + "messages": messages, + "stream": bool(body.get("stream")), + "n": 1, + } + for src, dst in ( + ("temperature", "temperature"), + ("top_p", "top_p"), + ("max_output_tokens", "max_completion_tokens"), + ("parallel_tool_calls", "parallel_tool_calls"), + ("user", "user"), + ): + if body.get(src) is not None: + chat[dst] = body[src] + + if body.get("tools"): + chat["tools"] = _tools_to_chat(body["tools"]) + tool_choice = body.get("tool_choice") + if tool_choice is not None: + if isinstance(tool_choice, dict) and tool_choice.get("type") == "function": + chat["tool_choice"] = {"type": "function", "function": {"name": tool_choice.get("name")}} + else: + chat["tool_choice"] = tool_choice + + effort = (body.get("reasoning") or {}).get("effort") + if effort in ("low", "medium", "high"): + chat["reasoning_effort"] = effort + + response_format = _text_format_to_response_format(body) + if response_format: + chat["response_format"] = response_format + + extra_body = body.get("extra_body") + if isinstance(extra_body, dict): + for k, v in extra_body.items(): + chat.setdefault(k, v) + + return chat + + +def _new_ids() -> Dict[str, str]: + return { + "response": f"resp_{uuid.uuid4().hex}", + "message": f"msg_{uuid.uuid4().hex}", + "reasoning": f"rs_{uuid.uuid4().hex}", + } + + +def _response_envelope(body: Dict[str, Any], response_id: str, created_at: int) -> Dict[str, Any]: + return { + "id": response_id, + "object": "response", + "created_at": created_at, + "status": "in_progress", + "background": False, + "error": None, + "incomplete_details": None, + "instructions": body.get("instructions"), + "max_output_tokens": body.get("max_output_tokens"), + "model": body.get("model", "default"), + "output": [], + "parallel_tool_calls": body.get("parallel_tool_calls", True), + "previous_response_id": None, + "reasoning": body.get("reasoning"), + "store": False, + "temperature": body.get("temperature"), + "text": body.get("text") or {"format": {"type": "text"}}, + "tool_choice": body.get("tool_choice", "auto"), + "tools": body.get("tools") or [], + "top_p": body.get("top_p"), + "truncation": body.get("truncation", "disabled"), + "usage": None, + "user": body.get("user"), + "metadata": body.get("metadata") or {}, + } + + +def _usage_to_responses(usage: Dict[str, Any]) -> Dict[str, Any]: + input_tokens = int(usage.get("prompt_tokens", 0)) + output_tokens = int(usage.get("completion_tokens", 0)) + cached = int((usage.get("prompt_tokens_details") or {}).get("cached_tokens", 0) or 0) + return { + "input_tokens": input_tokens, + "input_tokens_details": {"cached_tokens": cached}, + "output_tokens": output_tokens, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": input_tokens + output_tokens, + } + + +def _chat_response_to_responses(chat_response: Any, body: Dict[str, Any]) -> Dict[str, Any]: + if hasattr(chat_response, "model_dump"): + openai_dict = chat_response.model_dump(exclude_none=True) + else: + openai_dict = dict(chat_response) + + ids = _new_ids() + created_at = int(openai_dict.get("created") or time.time()) + result = _response_envelope(body, ids["response"], created_at) + + choice = (openai_dict.get("choices") or [{}])[0] + message = choice.get("message") or {} + finish_reason = choice.get("finish_reason") + + output: List[Dict[str, Any]] = [] + reasoning_text = message.get("reasoning") or message.get("reasoning_content") + if reasoning_text: + output.append( + { + "type": "reasoning", + "id": ids["reasoning"], + "summary": [], + "content": [{"type": "reasoning_text", "text": reasoning_text}], + } + ) + text = message.get("content") + if text: + output.append( + { + "type": "message", + "id": ids["message"], + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + ) + for tc in message.get("tool_calls") or []: + fn = tc.get("function") or {} + output.append( + { + "type": "function_call", + "id": f"fc_{uuid.uuid4().hex}", + "call_id": tc.get("id"), + "name": fn.get("name"), + "arguments": fn.get("arguments") or "", + "status": "completed", + } + ) + + result["output"] = output + result["usage"] = _usage_to_responses(openai_dict.get("usage") or {}) + if finish_reason == "length": + result["status"] = "incomplete" + result["incomplete_details"] = {"reason": "max_output_tokens"} + else: + result["status"] = "completed" + return result + + +async def _openai_sse_to_responses_events( + openai_body_iterator, + body: Dict[str, Any], +) -> AsyncGenerator[bytes, None]: + seq = 0 + + def event(event_type: str, data: Dict[str, Any]) -> bytes: + nonlocal seq + seq += 1 + data = {"type": event_type, "sequence_number": seq, **data} + return f"event: {event_type}\ndata: {json.dumps(data, ensure_ascii=False)}\n\n".encode("utf-8") + + ids = _new_ids() + response = _response_envelope(body, ids["response"], int(time.time())) + + yield event("response.created", {"response": response}) + yield event("response.in_progress", {"response": response}) + + output_index = -1 + current: Optional[tuple] = None + finish_reason = None + usage: Dict[str, Any] = {} + failed_error: Optional[Dict[str, Any]] = None + + def close_current(): + nonlocal current + if current is None: + return + kind, item = current + current = None + if kind == "message": + text = item["content"][0]["text"] + yield event( + "response.output_text.done", + {"item_id": item["id"], "output_index": output_index, "content_index": 0, "text": text}, + ) + yield event( + "response.content_part.done", + { + "item_id": item["id"], + "output_index": output_index, + "content_index": 0, + "part": {"type": "output_text", "text": text, "annotations": []}, + }, + ) + elif kind == "reasoning": + yield event( + "response.reasoning_text.done", + { + "item_id": item["id"], + "output_index": output_index, + "content_index": 0, + "text": item["content"][0]["text"], + }, + ) + elif kind == "function_call": + yield event( + "response.function_call_arguments.done", + {"item_id": item["id"], "output_index": output_index, "arguments": item["arguments"]}, + ) + item["status"] = "completed" + response["output"].append(item) + yield event("response.output_item.done", {"output_index": output_index, "item": item}) + + def open_item(kind: str, item: Dict[str, Any]): + nonlocal current, output_index + yield from close_current() + output_index += 1 + current = (kind, item) + yield event("response.output_item.added", {"output_index": output_index, "item": item}) + if kind == "message": + yield event( + "response.content_part.added", + { + "item_id": item["id"], + "output_index": output_index, + "content_index": 0, + "part": {"type": "output_text", "text": "", "annotations": []}, + }, + ) + + async for raw_chunk in openai_body_iterator: + if not raw_chunk: + continue + if isinstance(raw_chunk, (bytes, bytearray)): + raw_chunk = raw_chunk.decode("utf-8", errors="replace") + for line in raw_chunk.split("\n"): + line = line.strip() + if not line.startswith("data: "): + continue + payload = line[len("data: ") :] + if payload == "[DONE]": + continue + try: + chunk = json.loads(payload) + except Exception: + logger.debug("Skipping non-JSON SSE payload: %r", payload) + continue + + if "error" in chunk and "choices" not in chunk: + failed_error = chunk["error"] + break + + if chunk.get("usage"): + usage = chunk["usage"] + choices = chunk.get("choices") or [] + if not choices: + continue + choice = choices[0] + delta = choice.get("delta") or {} + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] + + reasoning_piece = delta.get("reasoning") or delta.get("reasoning_content") + if reasoning_piece: + if current is None or current[0] != "reasoning": + item = { + "type": "reasoning", + "id": f"rs_{uuid.uuid4().hex}", + "summary": [], + "content": [{"type": "reasoning_text", "text": ""}], + "status": "in_progress", + } + for e in open_item("reasoning", item): + yield e + item = current[1] + item["content"][0]["text"] += reasoning_piece + yield event( + "response.reasoning_text.delta", + { + "item_id": item["id"], + "output_index": output_index, + "content_index": 0, + "delta": reasoning_piece, + }, + ) + + content_piece = delta.get("content") + if content_piece: + if current is None or current[0] != "message": + item = { + "type": "message", + "id": f"msg_{uuid.uuid4().hex}", + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": "", "annotations": []}], + } + for e in open_item("message", item): + yield e + item = current[1] + item["content"][0]["text"] += content_piece + yield event( + "response.output_text.delta", + { + "item_id": item["id"], + "output_index": output_index, + "content_index": 0, + "delta": content_piece, + }, + ) + + for tc in delta.get("tool_calls") or []: + fn = tc.get("function") or {} + if fn.get("name"): + item = { + "type": "function_call", + "id": f"fc_{uuid.uuid4().hex}", + "call_id": tc.get("id") or f"call_{uuid.uuid4().hex[:24]}", + "name": fn["name"], + "arguments": "", + "status": "in_progress", + } + for e in open_item("function_call", item): + yield e + args = fn.get("arguments") + if args and current is not None and current[0] == "function_call": + item = current[1] + item["arguments"] += args + yield event( + "response.function_call_arguments.delta", + {"item_id": item["id"], "output_index": output_index, "delta": args}, + ) + if failed_error is not None: + break + + for e in close_current(): + yield e + + if failed_error is not None: + response["status"] = "failed" + response["error"] = {"code": "server_error", "message": failed_error.get("message", "generation failed")} + yield event("response.failed", {"response": response}) + return + + response["usage"] = _usage_to_responses(usage) + if finish_reason == "length": + response["status"] = "incomplete" + response["incomplete_details"] = {"reason": "max_output_tokens"} + yield event("response.incomplete", {"response": response}) + else: + response["status"] = "completed" + yield event("response.completed", {"response": response}) + + +async def responses_impl(raw_request: Request) -> Response: + from .api_models import ChatCompletionRequest, ChatCompletionResponse + from .api_openai import chat_completions_impl, create_error_response + + try: + body = await raw_request.json() + except Exception as exc: + return create_error_response(HTTPStatus.BAD_REQUEST, f"Invalid JSON body: {exc}") + if not isinstance(body, dict): + return create_error_response(HTTPStatus.BAD_REQUEST, "Request body must be a JSON object") + + if body.get("previous_response_id"): + return create_error_response( + HTTPStatus.BAD_REQUEST, + "previous_response_id is not supported (this server is stateless); " + "resend the full conversation in 'input' instead", + param="previous_response_id", + ) + if body.get("background"): + return create_error_response( + HTTPStatus.BAD_REQUEST, "background responses are not supported", param="background" + ) + + try: + chat_dict = _responses_to_chat_request(body) + chat_request = ChatCompletionRequest(**chat_dict) + except ValueError as exc: + return create_error_response(HTTPStatus.BAD_REQUEST, str(exc)) + except Exception as exc: + logger.exception("Failed to translate Responses API request") + return create_error_response(HTTPStatus.BAD_REQUEST, f"Invalid request: {exc}") + + downstream = await chat_completions_impl(chat_request, raw_request) + + if chat_request.stream: + if not isinstance(downstream, StreamingResponse): + return downstream + return StreamingResponse( + _openai_sse_to_responses_events(downstream.body_iterator, body), + media_type="text/event-stream", + ) + + if not isinstance(downstream, ChatCompletionResponse): + return downstream + + try: + return JSONResponse(_chat_response_to_responses(downstream, body)) + except Exception as exc: + logger.error("Failed to translate response to Responses API format: %s", exc, exc_info=True) + return create_error_response(HTTPStatus.INTERNAL_SERVER_ERROR, str(exc)) From 0f497621d5299123bef0033505d073159d43952a Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 17 Jul 2026 17:47:57 +0800 Subject: [PATCH 087/214] fix: align Responses API behavior --- lightllm/server/api_responses.py | 24 +++++-- unit_tests/server/test_api_responses.py | 87 +++++++++++++++++++++++++ 2 files changed, 105 insertions(+), 6 deletions(-) create mode 100644 unit_tests/server/test_api_responses.py diff --git a/lightllm/server/api_responses.py b/lightllm/server/api_responses.py index 4ce0e19673..c09dd3bdd6 100644 --- a/lightllm/server/api_responses.py +++ b/lightllm/server/api_responses.py @@ -128,6 +128,14 @@ def _text_format_to_response_format(body: Dict[str, Any]) -> Optional[Dict[str, def _responses_to_chat_request(body: Dict[str, Any]) -> Dict[str, Any]: + truncation = body.get("truncation") + if truncation not in (None, "disabled"): + raise ValueError("Only truncation='disabled' is supported") + + effort = (body.get("reasoning") or {}).get("effort") + if effort is not None and effort not in ("low", "medium", "high"): + raise ValueError("reasoning.effort must be one of: low, medium, high") + messages: List[Dict[str, Any]] = [] if body.get("instructions"): messages.append({"role": "system", "content": body["instructions"]}) @@ -177,8 +185,7 @@ def _responses_to_chat_request(body: Dict[str, Any]) -> Dict[str, Any]: else: chat["tool_choice"] = tool_choice - effort = (body.get("reasoning") or {}).get("effort") - if effort in ("low", "medium", "high"): + if effort is not None: chat["reasoning_effort"] = effort response_format = _text_format_to_response_format(body) @@ -360,7 +367,12 @@ def close_current(): elif kind == "function_call": yield event( "response.function_call_arguments.done", - {"item_id": item["id"], "output_index": output_index, "arguments": item["arguments"]}, + { + "item_id": item["id"], + "output_index": output_index, + "name": item["name"], + "arguments": item["arguments"], + }, ) item["status"] = "completed" response["output"].append(item) @@ -487,15 +499,15 @@ def open_item(kind: str, item: Dict[str, Any]): if failed_error is not None: break - for e in close_current(): - yield e - if failed_error is not None: response["status"] = "failed" response["error"] = {"code": "server_error", "message": failed_error.get("message", "generation failed")} yield event("response.failed", {"response": response}) return + for e in close_current(): + yield e + response["usage"] = _usage_to_responses(usage) if finish_reason == "length": response["status"] = "incomplete" diff --git a/unit_tests/server/test_api_responses.py b/unit_tests/server/test_api_responses.py new file mode 100644 index 0000000000..71750f06cf --- /dev/null +++ b/unit_tests/server/test_api_responses.py @@ -0,0 +1,87 @@ +import asyncio + +import pytest +import ujson as json + +from lightllm.server.api_responses import _openai_sse_to_responses_events, _responses_to_chat_request + + +async def _chunks(*payloads): + for payload in payloads: + yield f"data: {json.dumps(payload)}\n\n" + + +def _collect_events(*payloads): + async def collect(): + return [event async for event in _openai_sse_to_responses_events(_chunks(*payloads), {"input": "hi"})] + + raw_events = asyncio.run(collect()) + return [json.loads(event.decode().split("data: ", 1)[1]) for event in raw_events] + + +def test_function_call_arguments_done_includes_name(): + events = _collect_events( + { + "choices": [ + { + "delta": { + "tool_calls": [ + { + "id": "call_1", + "function": {"name": "get_weather", "arguments": '{"city":"Paris"}'}, + } + ] + } + } + ] + } + ) + + done = next(event for event in events if event["type"] == "response.function_call_arguments.done") + assert done["name"] == "get_weather" + + +@pytest.mark.parametrize( + "partial_payload", + [ + {"choices": [{"delta": {"content": "partial"}}]}, + { + "choices": [ + { + "delta": { + "tool_calls": [ + { + "id": "call_1", + "function": {"name": "get_weather", "arguments": '{"city":'}, + } + ] + } + } + ] + }, + ], +) +def test_stream_failure_does_not_complete_partial_item(partial_payload): + events = _collect_events(partial_payload, {"error": {"message": "generation failed"}}) + event_types = [event["type"] for event in events] + + assert "response.output_text.done" not in event_types + assert "response.function_call_arguments.done" not in event_types + assert "response.output_item.done" not in event_types + assert event_types[-1] == "response.failed" + + +@pytest.mark.parametrize("effort", ["none", "minimal", "xhigh"]) +def test_unsupported_reasoning_effort_is_rejected(effort): + with pytest.raises(ValueError, match="reasoning.effort"): + _responses_to_chat_request({"input": "hi", "reasoning": {"effort": effort}}) + + +def test_supported_reasoning_effort_is_forwarded(): + request = _responses_to_chat_request({"input": "hi", "reasoning": {"effort": "high"}}) + assert request["reasoning_effort"] == "high" + + +def test_automatic_truncation_is_rejected(): + with pytest.raises(ValueError, match="truncation"): + _responses_to_chat_request({"input": "hi", "truncation": "auto"}) From 046034436633a8a0ff9ada6f6a983493b5c573bf Mon Sep 17 00:00:00 2001 From: sufubao Date: Fri, 17 Jul 2026 18:24:19 +0800 Subject: [PATCH 088/214] fix: correct Responses API edge cases --- lightllm/server/api_http.py | 2 + lightllm/server/api_responses.py | 14 +-- unit_tests/server/test_api_responses.py | 109 +++++++++++++++++++++++- 3 files changed, 119 insertions(+), 6 deletions(-) diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index b433cf748d..eda572bd03 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -366,6 +366,8 @@ async def openai_responses(raw_request: Request) -> Response: except ServerBusyError as e: logger.error("%s", str(e), exc_info=True) return create_error_response(HTTPStatus.SERVICE_UNAVAILABLE, str(e)) + except ValueError as e: + return create_error_response(HTTPStatus.BAD_REQUEST, str(e)) except ClientDisconnected as e: logger.warning(str(e)) return Response(status_code=499) diff --git a/lightllm/server/api_responses.py b/lightllm/server/api_responses.py index c09dd3bdd6..a0ca3613de 100644 --- a/lightllm/server/api_responses.py +++ b/lightllm/server/api_responses.py @@ -65,7 +65,9 @@ def _input_items_to_messages(items: List[Any]) -> List[Dict[str, Any]]: ) elif itype == "function_call_output": output = item.get("output") - if not isinstance(output, str): + if isinstance(output, list): + output = _content_parts_to_chat(output) + elif not isinstance(output, str): output = json.dumps(output, ensure_ascii=False) messages.append({"role": "tool", "tool_call_id": item.get("call_id"), "content": output}) elif itype == "reasoning": @@ -263,6 +265,7 @@ def _chat_response_to_responses(chat_response: Any, body: Dict[str, Any]) -> Dic choice = (openai_dict.get("choices") or [{}])[0] message = choice.get("message") or {} finish_reason = choice.get("finish_reason") + output_status = "incomplete" if finish_reason == "length" else "completed" output: List[Dict[str, Any]] = [] reasoning_text = message.get("reasoning") or message.get("reasoning_content") @@ -281,7 +284,7 @@ def _chat_response_to_responses(chat_response: Any, body: Dict[str, Any]) -> Dic { "type": "message", "id": ids["message"], - "status": "completed", + "status": output_status, "role": "assistant", "content": [{"type": "output_text", "text": text, "annotations": []}], } @@ -295,7 +298,7 @@ def _chat_response_to_responses(chat_response: Any, body: Dict[str, Any]) -> Dic "call_id": tc.get("id"), "name": fn.get("name"), "arguments": fn.get("arguments") or "", - "status": "completed", + "status": output_status, } ) @@ -374,7 +377,7 @@ def close_current(): "arguments": item["arguments"], }, ) - item["status"] = "completed" + item["status"] = "incomplete" if finish_reason == "length" else "completed" response["output"].append(item) yield event("response.output_item.done", {"output_index": output_index, "item": item}) @@ -383,7 +386,8 @@ def open_item(kind: str, item: Dict[str, Any]): yield from close_current() output_index += 1 current = (kind, item) - yield event("response.output_item.added", {"output_index": output_index, "item": item}) + added_item = {**item, "content": []} if kind == "message" else item + yield event("response.output_item.added", {"output_index": output_index, "item": added_item}) if kind == "message": yield event( "response.content_part.added", diff --git a/unit_tests/server/test_api_responses.py b/unit_tests/server/test_api_responses.py index 71750f06cf..70048be5a6 100644 --- a/unit_tests/server/test_api_responses.py +++ b/unit_tests/server/test_api_responses.py @@ -3,7 +3,11 @@ import pytest import ujson as json -from lightllm.server.api_responses import _openai_sse_to_responses_events, _responses_to_chat_request +from lightllm.server.api_responses import ( + _chat_response_to_responses, + _openai_sse_to_responses_events, + _responses_to_chat_request, +) async def _chunks(*payloads): @@ -41,6 +45,109 @@ def test_function_call_arguments_done_includes_name(): assert done["name"] == "get_weather" +def test_content_array_function_output_is_preserved(): + request = _responses_to_chat_request( + { + "input": [ + { + "type": "function_call_output", + "call_id": "call_1", + "output": [ + {"type": "input_text", "text": "sunny"}, + {"type": "input_image", "image_url": "https://example.com/weather.png"}, + ], + } + ] + } + ) + + assert request["messages"] == [ + { + "role": "tool", + "tool_call_id": "call_1", + "content": [ + {"type": "text", "text": "sunny"}, + {"type": "image_url", "image_url": {"url": "https://example.com/weather.png"}}, + ], + } + ] + + +def test_non_streamed_truncated_function_call_is_incomplete(): + response = _chat_response_to_responses( + { + "choices": [ + { + "finish_reason": "length", + "message": { + "tool_calls": [ + { + "id": "call_1", + "function": {"name": "get_weather", "arguments": '{"city":'}, + } + ] + }, + } + ] + }, + {"input": "hi"}, + ) + + assert response["status"] == "incomplete" + assert response["output"][0]["status"] == "incomplete" + + +def test_streamed_truncated_function_call_is_incomplete(): + events = _collect_events( + { + "choices": [ + { + "delta": { + "tool_calls": [ + { + "id": "call_1", + "function": {"name": "get_weather", "arguments": '{"city":'}, + } + ] + } + } + ] + }, + {"choices": [{"delta": {}, "finish_reason": "length"}]}, + ) + + item_done = next(event for event in events if event["type"] == "response.output_item.done") + response_done = events[-1] + assert item_done["item"]["status"] == "incomplete" + assert response_done["type"] == "response.incomplete" + assert response_done["response"]["output"][0]["status"] == "incomplete" + + +def test_streamed_message_adds_text_content_once(): + events = _collect_events({"choices": [{"delta": {"content": "hello"}, "finish_reason": "stop"}]}) + + item_added = next(event for event in events if event["type"] == "response.output_item.added") + part_added = next(event for event in events if event["type"] == "response.content_part.added") + assert item_added["item"]["content"] == [] + assert part_added["part"] == {"type": "output_text", "text": "", "annotations": []} + + +def test_route_maps_downstream_value_error_to_bad_request(monkeypatch): + from types import SimpleNamespace + + from lightllm.server import api_http, api_responses + + async def raise_value_error(raw_request): + raise ValueError("Unrecognized image input.") + + monkeypatch.setattr(api_http, "get_env_start_args", lambda: SimpleNamespace(run_mode="normal")) + monkeypatch.setattr(api_http.g_objs, "metric_client", SimpleNamespace(counter_inc=lambda *args: None)) + monkeypatch.setattr(api_responses, "responses_impl", raise_value_error) + + response = asyncio.run(api_http.openai_responses(None)) + assert response.status_code == 400 + + @pytest.mark.parametrize( "partial_payload", [ From 77c8247c6ba3fcfd1af223ea00235664026e9acc Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 20 Jul 2026 04:19:28 +0000 Subject: [PATCH 089/214] use peak SWA demand for paused request recovery --- .../server/router/model_infer/infer_batch.py | 21 ++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 5e3de340cc..781beddcae 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -392,9 +392,7 @@ def recover_paused_reqs( or can_alloc_dsv4_c4_page_num is not None or can_alloc_dsv4_c128_slot_num is not None ): - swa_page_num, c4_page_num, c128_slot_num = req.get_dsv4_prefill_need_page_and_slot_num( - is_chuncked_prefill=False - ) + swa_page_num, c4_page_num, c128_slot_num = req.get_dsv4_recover_need_page_and_slot_num() if can_alloc_dsv4_swa_page_num is not None and swa_page_num > can_alloc_dsv4_swa_page_num: break if can_alloc_dsv4_c4_page_num is not None and c4_page_num > can_alloc_dsv4_c4_page_num: @@ -1018,6 +1016,23 @@ def get_dsv4_prefill_need_page_and_slot_num(self, is_chuncked_prefill: bool) -> c128_slot_num = max(0, end // 128 - start // 128) if self.dsv4_has_c128 else 0 return swa_page_num, c4_page_num, c128_slot_num + def get_dsv4_recover_need_page_and_slot_num(self) -> Tuple[int, int, int]: + swa_page_num, c4_page_num, c128_slot_num = self.get_dsv4_prefill_need_page_and_slot_num( + is_chuncked_prefill=False + ) + if swa_page_num == 0 or self.args.disable_chunked_prefill: + return swa_page_num, c4_page_num, c128_slot_num + + # C4/C128 accumulate across recovery chunks; only SWA is evicted chunk by chunk. + req_manager: DeepseekV4ReqManager = g_infer_context.req_manager + prompt_cache_page_size = req_manager.get_prompt_cache_page_size() + peak_token_num = min( + self.get_cur_total_len(), + self.args.chunked_prefill_size + int(req_manager.sliding_window) + 2 * prompt_cache_page_size, + ) + swa_page_num = (peak_token_num + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size + return swa_page_num, c4_page_num, c128_slot_num + def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: seq_len = self.get_cur_total_len() if seq_len <= 0: From 5925ca66d9c2027c259d319a40000d23182838e1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 21 Jul 2026 15:14:59 +0000 Subject: [PATCH 090/214] DeepSeek DSML parsing --- lightllm/server/function_call_parser.py | 72 +++++++++++++------------ 1 file changed, 37 insertions(+), 35 deletions(-) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index a66e383a46..26d2dbe45e 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1532,30 +1532,30 @@ def _dsml_params_to_json(self, params: List[tuple]) -> str: def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult: """One-time parsing for DSML format tool calls.""" - idx = text.find(self.bot_token) - normal_text = text[:idx].strip() if idx != -1 else text - if self.bot_token not in text: - return StreamingParseResult(normal_text=normal_text, calls=[]) + block_start = text.find(self.bot_token) + if block_start == -1: + return StreamingParseResult(normal_text=text, calls=[]) + + block_body_start = block_start + len(self.bot_token) + block_end = text.find(self.eot_token, block_body_start) + if block_end == -1: + return StreamingParseResult(normal_text=text, calls=[]) + + normal_text = text[:block_start].removesuffix("\n\n") - tool_indices = self._get_tool_indices(tools) calls = [] - invoke_matches = self.invoke_regex.findall(text) + invoke_matches = self.invoke_regex.findall(text[block_body_start:block_end]) for func_name, invoke_body in invoke_matches: - if func_name not in tool_indices: - logger.warning(f"Model attempted to call undefined function: {func_name}") - continue - param_matches = self.param_regex.findall(invoke_body) args_json = self._dsml_params_to_json(param_matches) - - calls.append( - ToolCallItem( - tool_index=tool_indices[func_name], - name=func_name, - parameters=args_json, - ) - ) + match_result = { + "name": func_name, + "parameters": json.loads(args_json), + } + for item in self.parse_base_json(match_result, tools): + item.tool_index = len(calls) + calls.append(item) return StreamingParseResult(normal_text=normal_text, calls=calls) @@ -1579,17 +1579,20 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami normal_text = normal_text.replace(e_token, "") return StreamingParseResult(normal_text=normal_text) + normal_text = "" + # Mark that we're inside a function_calls block if self.has_tool_call(current_text): + block_start = current_text.find(self.bot_token) + normal_text = current_text[:block_start].removesuffix("\n\n") + current_text = current_text[block_start:] + self._buffer = current_text self._in_function_calls = True # Check if function_calls block has ended if self.eot_token in current_text: self._in_function_calls = False - if not hasattr(self, "_tool_indices"): - self._tool_indices = self._get_tool_indices(tools) - calls: List[ToolCallItem] = [] try: @@ -1681,19 +1684,18 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami self.streamed_args_for_tool.append("") if not self.current_tool_name_sent: - if func_name in self._tool_indices: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=func_name, - parameters="", - ) + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + name=func_name, + parameters="", ) - self.current_tool_name_sent = True - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": {}, - } + ) + self.current_tool_name_sent = True + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": {}, + } else: # Stream arguments as complete parameters are parsed param_matches = self.param_regex.findall(partial_body) @@ -1720,11 +1722,11 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami except json.JSONDecodeError: pass - return StreamingParseResult(normal_text="", calls=calls) + return StreamingParseResult(normal_text=normal_text, calls=calls) except Exception as e: logger.error(f"Error in DeepSeekV32 parse_streaming_increment: {e}") - return StreamingParseResult(normal_text="", calls=calls) + return StreamingParseResult(normal_text=normal_text, calls=calls) class Qwen3CoderDetector(BaseFormatDetector): From d67f54b03a22b702a5ef7b648ca96075062e574f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 22 Jul 2026 09:04:50 +0000 Subject: [PATCH 091/214] Deleted some useless --- .../attention/nsa/dsv4_fp8_flashmla_sparse.py | 22 ++++--------------- lightllm/models/deepseek_v4/infer_struct.py | 8 +++---- .../deepseek_v4/layer_infer/compressor.py | 19 ++++------------ .../layer_infer/hyper_connection.py | 2 +- 4 files changed, 13 insertions(+), 38 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py index eece2648fb..52e38fd4f0 100644 --- a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING import torch +from vllm.v1.attention.ops import flashmla from ..base_att import AttControl, BaseAttBackend, BaseDecodeAttState, BasePrefillAttState @@ -27,12 +28,6 @@ def _view_cache(buffer: torch.Tensor, page_size: int) -> torch.Tensor: return buffer[:, :byte_num].view(buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) -def _get_vllm_flashmla(): - from vllm.v1.attention.ops import flashmla - - return flashmla - - class DeepseekV4FlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) @@ -47,7 +42,6 @@ def _flashmla_att( mem_manager, nsa_dict: dict, sched_meta, - flash_mla, flashmla_out: torch.Tensor = None, ) -> torch.Tensor: from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( @@ -88,7 +82,7 @@ def _flashmla_att( ) if flashmla_out is not None: kwargs["out"] = flashmla_out - full_out, _ = flash_mla.flash_mla_with_kvcache(**kwargs) + full_out, _ = flashmla.flash_mla_with_kvcache(**kwargs) return full_out[:, 0, : self.real_q_head_num, :] def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "_PrefillAttState": @@ -101,15 +95,13 @@ def create_att_decode_state(self, infer_state: "InferStateInfo") -> "_DecodeAttS @dataclasses.dataclass class _PrefillAttState(BasePrefillAttState): flashmla_sched_meta: dict = None - flash_mla: object = None def init_state(self): - self.flash_mla = _get_vllm_flashmla() self.flashmla_sched_meta = {} def _get_sched_meta(self, compress_ratio: int): if compress_ratio not in self.flashmla_sched_meta: - self.flashmla_sched_meta[compress_ratio] = self.flash_mla.get_mla_metadata()[0] + self.flashmla_sched_meta[compress_ratio] = flashmla.get_mla_metadata()[0] return self.flashmla_sched_meta[compress_ratio] def prefill_att( @@ -139,7 +131,6 @@ def prefill_att( self.infer_state.mem_manager, nsa_dict, self._get_sched_meta(nsa_dict["compress_ratio"]), - self.flash_mla, flashmla_out=full_out, ) ) @@ -149,17 +140,13 @@ def prefill_att( @dataclasses.dataclass class _DecodeAttState(BaseDecodeAttState): flashmla_sched_meta: dict = None - flash_mla: object = None def init_state(self): self.reset_sched_meta_for_capture() def reset_sched_meta_for_capture(self): # FlashMLA lazily binds extra-cache geometry, so ratios cannot share one sched-meta object. - self.flash_mla = _get_vllm_flashmla() - self.flashmla_sched_meta = { - ratio: self.flash_mla.get_mla_metadata()[0] for ratio in self.backend.compress_ratios - } + self.flashmla_sched_meta = {ratio: flashmla.get_mla_metadata()[0] for ratio in self.backend.compress_ratios} def decode_att( self, @@ -178,7 +165,6 @@ def decode_att( self.infer_state.mem_manager, nsa_dict, self.flashmla_sched_meta[nsa_dict["compress_ratio"]], - self.flash_mla, ) return real_out.contiguous() diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index e14e84c968..488cf82f92 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -53,12 +53,12 @@ def init_some_extra_state(self, model): # Per-token request id (decode: one token per req; prefill: ragged -> repeat by q-len). # Layer-independent; the swa kernel + build_metadata's c4/c128 readers all reuse it. if self.is_prefill: - self.dsv4_sparse_req_idx = torch.repeat_interleave(self.b_req_idx, self.b_q_seq_len.long()) + self.dsv4_sparse_req_idx = torch.repeat_interleave(self.b_req_idx, self.b_q_seq_len) self._dsv4_token_to_batch_idx = torch.repeat_interleave( - torch.arange(self.b_req_idx.shape[0], device=self.b_req_idx.device), - self.b_q_seq_len.long(), + torch.arange(self.b_req_idx.shape[0], dtype=torch.int32, device=self.b_req_idx.device), + self.b_q_seq_len, output_size=pos.numel(), - ).to(torch.int32) + ) else: self.dsv4_sparse_req_idx = self.b_req_idx self._dsv4_token_to_batch_idx = None diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 25e1df8601..149abfd360 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -302,7 +302,7 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, # out_buffer (the kernel's OUTPUT_BF16 path); the fp8 pack into c4_indexer_pool is done # afterwards by pack_indexer_k_to_cache. assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) state_ring = mem_manager.c4_state_ring out_buffer = torch.empty( @@ -313,30 +313,19 @@ def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, out_page_size = 1 # unused under OUTPUT_BF16 (token-indexed dense scratch, not paged) else: if compress_ratio == 4: - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] state_buffer = mem_manager.get_c4_state_buffer(layer_idx) state_ring = mem_manager.c4_state_ring out_pool = mem_manager.c4_pool elif compress_ratio == 128: - out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.long().reshape(-1)] + out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.reshape(-1)] state_buffer = mem_manager.get_c128_state_buffer(layer_idx) state_ring = mem_manager.c128_state_ring out_pool = mem_manager.c128_pool - else: - raise AssertionError(f"invalid DeepSeek-V4 compress ratio {compress_ratio}") out_buffer = mem_manager.get_compressed_kv_buffer(layer_idx) out_page_size = out_pool.page_size - token_to_batch_idx = infer_state.b_req_idx - if infer_state.is_prefill: - token_to_batch_idx = getattr(infer_state, "_dsv4_token_to_batch_idx", None) - if token_to_batch_idx is None or token_to_batch_idx.numel() != infer_state.position_ids.numel(): - q_lens = (infer_state.b_seq_len - infer_state.b_ready_cache_len).to(torch.long) - batch_idx = torch.arange(infer_state.b_req_idx.shape[0], device=infer_state.b_req_idx.device) - token_to_batch_idx = torch.repeat_interleave( - batch_idx, q_lens, output_size=infer_state.position_ids.numel() - ).to(torch.int32) - infer_state._dsv4_token_to_batch_idx = token_to_batch_idx + token_to_batch_idx = infer_state._dsv4_token_to_batch_idx if infer_state.is_prefill else infer_state.b_req_idx return CoreCompressorMetadata( layer_idx=layer_idx, diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index 080ebabd89..c2592c4d0b 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -62,7 +62,7 @@ def hc_head(streams, hc_fn, hc_scale, hc_base, hc_mult, dim, rms_eps, hc_eps, al """Final stream collapse before the lm_head. streams:[N, hc*dim] -> [N, dim].""" out = alloc_func((streams.shape[0], dim), dtype=streams.dtype, device=streams.device) torch.ops.vllm.hc_head_fused_kernel_tilelang( - streams.view(-1, hc_mult, dim).contiguous(), + streams.view(-1, hc_mult, dim), hc_fn, hc_scale, hc_base, From ec1127525eb181d50496bc101223cbb540234c74 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 22 Jul 2026 13:47:18 +0000 Subject: [PATCH 092/214] delete metadata --- .../deepseek_v4/layer_infer/compressor.py | 184 ++++++------------ .../layer_infer/transformer_layer_infer.py | 123 ++++-------- 2 files changed, 104 insertions(+), 203 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 149abfd360..a270a8ac5f 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -1,38 +1,11 @@ -from dataclasses import dataclass -from typing import Optional - import torch import triton import triton.language as tl from triton.language.extra import libdevice -from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE -@dataclass -class CoreCompressorMetadata: - layer_idx: int - compress_ratio: int - out_slots: torch.Tensor - mem_index: torch.Tensor - state_buffer: torch.Tensor - state_ring: int - out_buffer: torch.Tensor - out_page_size: int - position_ids: torch.Tensor - b_req_idx: torch.Tensor - b_mtp_index: torch.Tensor - b_seq_len: torch.Tensor - b_ready_cache_len: Optional[torch.Tensor] - b_q_start_loc: Optional[torch.Tensor] - req_to_token_indexs: torch.Tensor - full_to_swa_indexs: torch.Tensor - token_to_batch_idx: Optional[torch.Tensor] - kv_score: Optional[torch.Tensor] - is_prefill: bool - - @triton.jit def _add_ape_to_kv_score_kernel( kv_score, @@ -291,73 +264,14 @@ def _fused_compress_norm_rope_insert_kernel( return -def prepare_compress_states(*, infer_state, layer_idx: int, compress_ratio: int, is_in_indexer: bool = False): - if compress_ratio == 0: - return None - - mem_manager: DeepseekV4MemoryManager = infer_state.mem_manager - if is_in_indexer: - # c4 Lightning-Indexer key compression: same window/state machinery as the c4 latent - # compressor but with index_head_dim, a separate state pool, and a DENSE bf16 scratch - # out_buffer (the kernel's OUTPUT_BF16 path); the fp8 pack into c4_indexer_pool is done - # afterwards by pack_indexer_k_to_cache. - assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] - state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) - state_ring = mem_manager.c4_state_ring - out_buffer = torch.empty( - (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), - dtype=torch.bfloat16, - device=infer_state.mem_index.device, - ) - out_page_size = 1 # unused under OUTPUT_BF16 (token-indexed dense scratch, not paged) - else: - if compress_ratio == 4: - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] - state_buffer = mem_manager.get_c4_state_buffer(layer_idx) - state_ring = mem_manager.c4_state_ring - out_pool = mem_manager.c4_pool - elif compress_ratio == 128: - out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.reshape(-1)] - state_buffer = mem_manager.get_c128_state_buffer(layer_idx) - state_ring = mem_manager.c128_state_ring - out_pool = mem_manager.c128_pool - out_buffer = mem_manager.get_compressed_kv_buffer(layer_idx) - out_page_size = out_pool.page_size - - token_to_batch_idx = infer_state._dsv4_token_to_batch_idx if infer_state.is_prefill else infer_state.b_req_idx - - return CoreCompressorMetadata( - layer_idx=layer_idx, - compress_ratio=compress_ratio, - out_slots=out_slots, - mem_index=infer_state.mem_index, - state_buffer=state_buffer, - state_ring=state_ring, - out_buffer=out_buffer, - out_page_size=out_page_size, - position_ids=infer_state.position_ids, - b_req_idx=infer_state.b_req_idx, - b_mtp_index=infer_state.b_mtp_index, - b_seq_len=infer_state.b_seq_len, - b_ready_cache_len=infer_state.b_ready_cache_len, - b_q_start_loc=infer_state.b_q_start_loc, - req_to_token_indexs=infer_state.req_manager.req_to_token_indexs, - full_to_swa_indexs=mem_manager.full_to_swa_indexs, - token_to_batch_idx=token_to_batch_idx, - kv_score=None, - is_prefill=infer_state.is_prefill, - ) - - def prepare_partial_states( *, kv_score: torch.Tensor, - metadata: Optional[CoreCompressorMetadata], + position_ids: torch.Tensor, ape: torch.Tensor, compress_ratio: int, ): - if metadata is None or kv_score.shape[0] == 0: + if kv_score.shape[0] == 0: return state_width = kv_score.shape[-1] // 2 _add_ape_to_kv_score_kernel[(kv_score.shape[0],)]( @@ -366,7 +280,7 @@ def prepare_partial_states( kv_score.stride(1), ape, ape.stride(0), - metadata.position_ids, + position_ids, STATE_WIDTH=state_width, COMPRESS_RATIO=compress_ratio, BLOCK=triton.next_power_of_2(state_width), @@ -378,9 +292,9 @@ def prepare_partial_states( def fused_compress( *, kv_score: torch.Tensor, - metadata: Optional[CoreCompressorMetadata], + infer_state, + layer_idx: int, norm_weight: torch.Tensor, - ape: torch.Tensor, eps: float, head_dim: int, qk_rope_head_dim: int, @@ -389,32 +303,62 @@ def fused_compress( sin_table: torch.Tensor, output_bf16: bool = False, ): - if metadata is None or kv_score.shape[0] == 0: - return + mem_manager = infer_state.mem_manager + if output_bf16: + assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) + state_ring = mem_manager.c4_state_ring + out_buffer = torch.empty( + (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), + dtype=torch.bfloat16, + device=infer_state.mem_index.device, + ) + out_page_size = 1 + else: + if compress_ratio == 4: + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c4_state_buffer(layer_idx) + state_ring = mem_manager.c4_state_ring + out_page_size = mem_manager.c4_pool.page_size + elif compress_ratio == 128: + out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c128_state_buffer(layer_idx) + state_ring = mem_manager.c128_state_ring + out_page_size = mem_manager.c128_pool.page_size + out_buffer = mem_manager.get_compressed_kv_buffer(layer_idx) + + if kv_score.shape[0] == 0: + return out_buffer state_width = kv_score.shape[-1] // 2 - state_last_dim = metadata.state_buffer.shape[-1] + state_last_dim = state_buffer.shape[-1] is_c4 = compress_ratio == 4 - state_ring = metadata.state_ring block_state = triton.next_power_of_2(state_width) block_head = triton.next_power_of_2(head_dim) + token_to_batch_idx = infer_state._dsv4_token_to_batch_idx if infer_state.is_prefill else infer_state.b_req_idx + ready_cache_len = ( + infer_state.b_ready_cache_len if infer_state.b_ready_cache_len is not None else infer_state.b_seq_len + ) + q_start_loc = infer_state.b_q_start_loc if infer_state.b_q_start_loc is not None else infer_state.b_seq_len + req_to_token_indexs = infer_state.req_manager.req_to_token_indexs _fused_compress_norm_rope_insert_kernel[(kv_score.shape[0],)]( kv_score, kv_score.stride(0), kv_score.stride(1), - metadata.state_buffer, - metadata.position_ids, - metadata.token_to_batch_idx, - metadata.b_req_idx, - metadata.b_mtp_index, - metadata.b_seq_len, - metadata.b_ready_cache_len if metadata.b_ready_cache_len is not None else metadata.b_seq_len, - metadata.b_q_start_loc if metadata.b_q_start_loc is not None else metadata.b_seq_len, - metadata.req_to_token_indexs, - metadata.req_to_token_indexs.stride(0), - metadata.full_to_swa_indexs, - metadata.out_slots, + state_buffer, + infer_state.position_ids, + token_to_batch_idx, + infer_state.b_req_idx, + infer_state.b_mtp_index, + infer_state.b_seq_len, + ready_cache_len, + q_start_loc, + req_to_token_indexs, + req_to_token_indexs.stride(0), + mem_manager.full_to_swa_indexs, + out_slots, norm_weight, eps, cos_table, @@ -423,14 +367,14 @@ def fused_compress( sin_table, sin_table.stride(0), sin_table.stride(1), - metadata.out_buffer, + out_buffer, HEAD_DIM=head_dim, STATE_WIDTH=state_width, STATE_LAST_DIM=state_last_dim, COMPRESS_RATIO=compress_ratio, WINDOW_SIZE=compress_ratio * (2 if is_c4 else 1), IS_C4=is_c4, - IS_PREFILL=metadata.is_prefill, + IS_PREFILL=infer_state.is_prefill, SWA_PAGE_SIZE=DSV4_SWA_PAGE_SIZE, STATE_RING=state_ring, ROPE_HEAD_DIM=qk_rope_head_dim, @@ -439,8 +383,8 @@ def fused_compress( NOPE_DIM=head_dim - qk_rope_head_dim, QUANT_BLOCK=64, SCALE_BYTES=(head_dim - qk_rope_head_dim) // 64 + 1, - PAGE_SIZE=metadata.out_page_size, - BYTES_PER_PAGE=metadata.out_buffer.shape[-1], + PAGE_SIZE=out_page_size, + BYTES_PER_PAGE=out_buffer.shape[-1], BLOCK=block_head, OUTPUT_BF16=output_bf16, num_warps=4, @@ -450,21 +394,21 @@ def fused_compress( kv_score, kv_score.stride(0), kv_score.stride(1), - metadata.position_ids, - metadata.token_to_batch_idx, - metadata.b_req_idx, - metadata.b_seq_len, - metadata.mem_index, - metadata.full_to_swa_indexs, - metadata.state_buffer, + infer_state.position_ids, + token_to_batch_idx, + infer_state.b_req_idx, + infer_state.b_seq_len, + infer_state.mem_index, + mem_manager.full_to_swa_indexs, + state_buffer, STATE_WIDTH=state_width, STATE_LAST_DIM=state_last_dim, COMPRESS_RATIO=compress_ratio, IS_C4=is_c4, - IS_PREFILL=metadata.is_prefill, + IS_PREFILL=infer_state.is_prefill, SWA_PAGE_SIZE=DSV4_SWA_PAGE_SIZE, STATE_RING=state_ring, BLOCK=block_state, num_warps=4, ) - return + return out_buffer diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 638e11f6f2..2bc482955f 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -12,7 +12,6 @@ from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post from .compressor import fused_compress as fused_compress_op from .compressor import prepare_partial_states -from .compressor import prepare_compress_states from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo import deep_gemm @@ -156,16 +155,11 @@ def _get_qkv( input = self._tpsp_allgather(input=input, infer_state=infer_state) T = input.shape[0] - # wq_a and wkv share `input` -> one fused fp8 GEMM, split [q_lora_rank | head_dim]. qa is a - # row-strided view (rmsnorm honors stride(0)); kv feeds the fused cache writer -> contiguous. + qkv = layer_weight.wq_a_wkv_.mm(input) qa = layer_weight.q_norm_(qkv[:, : -self.head_dim_], eps=self.eps_) q_in = layer_weight.wq_b_.mm(qa).view(T, self.tp_q_head_num_, self.head_dim_) - # per-(token, head) weightless self-RMSNorm + interleaved rope on the last rope_dim dims, - # fused in one DSV4 CUDA kernel (fp32 norm/rotation, bf16 in between -- same as eager). - # The selected FlashMLA MODEL1 binary only instantiates H=64/128. Produce that ABI layout - # directly. Prefill uses one max-token workspace so changing T only changes the prefix view; - # decode keeps its graph-owned tensor. The workspace's padded tail is zeroed once at init. + if infer_state.is_prefill: q = infer_state.dsv4_workspace.flashmla_prefill_q[:T] else: @@ -173,12 +167,10 @@ def _get_qkv( q[:, self.tp_q_head_num_ :, :].zero_() fused_q_norm_rope(q_in, q[:, : self.tp_q_head_num_, :], self.eps_, self.freqs_cis, infer_state.position_ids) # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 DSV4 CUDA kernel 完成, - # 替代 eager norm/rope/cat + _post_cache_kv。 - # bf16 kv 中间量没有其他消费者: flashmla 路径注意力读 cache,压缩器/indexer 取 x。 infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, mem_index=infer_state.mem_index, - kv=qkv[:, -self.head_dim_ :].contiguous(), + kv=qkv[:, -self.head_dim_ :], kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, freqs_cis=self.freqs_cis, @@ -215,9 +207,6 @@ def _context_attention_wrapper_run( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): if torch.cuda.is_current_stream_capturing(): - q = q.contiguous() - q_lora = q_lora.contiguous() - x = x.contiguous() _q = tensor_to_no_ref_tensor(q) _q_lora = tensor_to_no_ref_tensor(q_lora) _x = tensor_to_no_ref_tensor(x) @@ -245,10 +234,8 @@ def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, cos_table, sin_table = self.cos_compress_table, self.sin_compress_table aux_stream = self.dsv4_prefill_aux_stream if self.compress_ratio == 4 and aux_stream is not None and not torch.cuda.is_current_stream_capturing(): - # _dsv4_token_to_batch_idx is built in init_some_extra_state (default stream, before this fork), - # so both the aux indexer-compressor and the main compressor read a ready, race-free tensor. main_stream = torch.cuda.current_stream() - aux_stream.wait_stream(main_stream) # fork: aux waits for x / q_lora produced on main + aux_stream.wait_stream(main_stream) # aux waits for x / q_lora produced on main with torch.cuda.stream(aux_stream): # x / q_lora are main-allocated and read here -> record so the allocator won't reuse them. x.record_stream(aux_stream) @@ -259,8 +246,7 @@ def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, meta = self.index_infer.build_metadata( x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=False ) - self.compressor.prepare_states(x, infer_state, layer_weight) - self.compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) + self.compressor.compress(x, infer_state, layer_weight, cos_table, sin_table) main_stream.wait_stream(aux_stream) # join before prefill_att reads the indices / latent KV # extra_indices / extra_lengths were allocated on aux -> record on main so they survive until consumed. for _t in (meta.get("extra_indices"), meta.get("extra_lengths")): @@ -268,9 +254,7 @@ def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, _t.record_stream(main_stream) return meta - # serial fallback -- semantics identical to the original sequence. - self.compressor.prepare_states(x, infer_state, layer_weight) - self.compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) + self.compressor.compress(x, infer_state, layer_weight, cos_table, sin_table) # write c4 Lightning-Indexer keys BEFORE build_metadata so the scorer reads fresh+accumulated entries. self.index_infer.write_indexer_k(x, infer_state, layer_weight, cos_table, sin_table) return self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) @@ -321,8 +305,7 @@ def token_attention_forward( def _token_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - self.compressor.prepare_states(x, infer_state, layer_weight) - self.compressor.fused_compress(infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) + self.compressor.compress(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) att_control = AttControl( @@ -407,10 +390,8 @@ def _select_experts( indices_dtype = torch.int64 if self.is_hash: hash_indices_table = layer_weight.gate_tid2eid_.weight - if not hash_indices_table.is_contiguous(): - hash_indices_table = hash_indices_table.contiguous() indices_dtype = hash_indices_table.dtype - input_tokens = infer_state.input_ids.to(dtype=indices_dtype).contiguous() + input_tokens = infer_state.input_ids.to(dtype=indices_dtype) else: bias = layer_weight.gate_bias_.weight @@ -428,7 +409,7 @@ def _select_experts( input_tokens, hash_indices_table, ) - return weights, indices.long() + return weights, indices class CompressorInfer: @@ -449,69 +430,43 @@ def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int, is_ self.index_head_dim = network_config["index_head_dim"] self.qk_rope_head_dim = network_config["qk_rope_head_dim"] self.eps = network_config["rms_norm_eps"] - self._metadata = None - def prepare_states( + def compress( self, x: torch.Tensor, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, - use_custom_tensor_manager: bool = True, - ): - # use_custom_tensor_manager=False routes the .mm outputs through torch.empty (stream-aware) - # instead of the stream-blind global cache -- required when this runs on the prefill aux stream. - self._metadata = prepare_compress_states( - infer_state=infer_state, - layer_idx=self.layer_idx_, - compress_ratio=self.compress_ratio, - is_in_indexer=self.is_in_indexer, - ) - if self._metadata is not None: - if self.is_in_indexer: - # fused wkv/wgate GEMM -> [T, 2*coff*idx_hd] in the [kv | score] layout directly - # (same as the attention compressor_wkv_gate_). - self._metadata.kv_score = layer_weight.idx_cmp_wkv_gate_.mm( - x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 - ) - ape = layer_weight.idx_cmp_ape_.weight - else: - self._metadata.kv_score = layer_weight.compressor_wkv_gate_.mm( - x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 - ) - ape = layer_weight.compressor_ape_.weight - prepare_partial_states( - kv_score=self._metadata.kv_score, - metadata=self._metadata, - ape=ape, - compress_ratio=self.compress_ratio, - ) - return self._metadata - - def fused_compress( - self, - infer_state: DeepseekV4InferStateInfo, - layer_weight: DeepseekV4TransformerLayerWeight, cos_table: torch.Tensor, sin_table: torch.Tensor, + use_custom_tensor_manager: bool = True, ): if self.compress_ratio == 0: return None - metadata = self._metadata - if metadata is None: - raise RuntimeError("DeepSeek-V4 compressor.prepare_states must run before fused_compress") if self.is_in_indexer: + kv_score = layer_weight.idx_cmp_wkv_gate_.mm( + x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 + ) norm_weight = layer_weight.idx_cmp_norm_.weight ape = layer_weight.idx_cmp_ape_.weight head_dim = self.index_head_dim else: + kv_score = layer_weight.compressor_wkv_gate_.mm( + x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 + ) norm_weight = layer_weight.compressor_norm_.weight ape = layer_weight.compressor_ape_.weight head_dim = self.head_dim + prepare_partial_states( + kv_score=kv_score, + position_ids=infer_state.position_ids, + ape=ape, + compress_ratio=self.compress_ratio, + ) return fused_compress_op( - kv_score=metadata.kv_score, - metadata=metadata, + kv_score=kv_score, + infer_state=infer_state, + layer_idx=self.layer_idx_, norm_weight=norm_weight, - ape=ape, eps=self.eps, head_dim=head_dim, qk_rope_head_dim=self.qk_rope_head_dim, @@ -532,7 +487,7 @@ class DeepseekV4IndexInfer: init_some_extra_state; this class owns the c4 entry gather (build_compress_index) AND the c4 Lightning-Indexer scoring (gather + deep_gemm.fp8_mqa_logits + topk). Holds only static per-layer config; all per-request data flows in via args. Invoke from _context/_token_attention_kernel - (after compressor.fused_compress, before *_att) so the c4 scorer/topk keep the same cuda-graph + (after compressor.compress, before *_att) so the c4 scorer/topk keep the same cuda-graph capture position they had when this lived in the backend. The indexer is replicated (no TP collective).""" def __init__(self, layer_idx: int, network_config: dict, tp_world_size: int): @@ -568,11 +523,15 @@ def write_indexer_k( later long-context scoring. No-op on c128 / dense layers.""" if self.compress_ratio != 4: return - self.indexer_compressor.prepare_states( - x, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager + # Only group-end rows in this dense bf16 scratch are valid indexer keys. + scratch = self.indexer_compressor.compress( + x, + infer_state, + layer_weight, + cos_table, + sin_table, + use_custom_tensor_manager=use_custom_tensor_manager, ) - self.indexer_compressor.fused_compress(infer_state, layer_weight, cos_table, sin_table) - scratch = self.indexer_compressor._metadata.out_buffer # [T, index_head_dim] bf16 (group-end rows valid) # Rotate K (post norm+rope) by the SAME 1/sqrt(d) Hadamard the q kernel applies, so # (Hq)·(Hk)=q·k (H orthogonal) and the fp8 quant of K stays accurate. from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform @@ -580,7 +539,7 @@ def write_indexer_k( scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5) mem_manager = infer_state.mem_manager positions = infer_state.position_ids - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.long().reshape(-1)] + out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] # only group-end tokens finish a c4 entry; mask the rest to -1 so the packer skips them # (mid-group tokens share the group's c4 slot -> avoids racing a finished slot). completed = ((positions + 1) % 4 == 0) & (out_slots >= 0) @@ -646,7 +605,7 @@ def _indexer_q_weight( self.freqs_cis, infer_state.position_ids, ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded - return idx_q_fp8, weights.squeeze(-1).contiguous() + return idx_q_fp8, weights.squeeze(-1) def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): """c4 scorer via the page-safe deep_gemm.fp8_paged_mqa_logits over the paged c4 indexer pool, @@ -708,20 +667,18 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, infer_state.b_q_seq_len, output_size=positions.numel(), ) - row_page_table = page_table[token_batch_pos.long()].contiguous() + row_page_table = page_table[token_batch_pos] else: row_page_table = page_table valid_len = ((positions + 1) // 4).to(torch.int32) - ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1).contiguous() + ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1) metadata = deep_gemm.get_paged_mqa_logits_metadata( ctx_lens, page_size, deep_gemm.get_num_sms(), ) - topk_lengths = torch.clamp( - torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1 - ).contiguous() + topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) cached = (row_page_table, valid_len, ctx_lens, metadata, topk_lengths) infer_state._c4_paged_meta = cached From 99cf2a255976cad311ef64eef0b3b089c5a6acd5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 23 Jul 2026 06:20:27 +0000 Subject: [PATCH 093/214] Optimize DSV4 KV cache write and fix FP8 scaling --- .../deepseek4_mem_manager.py | 34 ++++++----- .../deepseek_v4/layer_infer/compressor.py | 58 ++++++++++++------- .../layer_infer/transformer_layer_infer.py | 27 +++------ .../layer_weights/transformer_layer_weight.py | 2 +- .../destindex_copy_indexer_k_dsv4.py | 39 +++++++++---- .../destindex_copy_kv_flashmla_dsv4.py | 8 +-- 6 files changed, 97 insertions(+), 71 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index abc15ecd0c..a5092cb08d 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -25,7 +25,7 @@ DSV4_INDEXER_SCALE_BYTES = 4 # 4B fp32 scale DSV4_INDEXER_BYTES_PER_TOKEN = DSV4_INDEXER_HEAD_DIM + DSV4_INDEXER_SCALE_BYTES # 132 DSV4_FP8_E4M3_MAX = 448.0 # 448.0 -DSV4_FP8_SCALE_MIN = 1e-4 # 1e-4 +DSV4_FP8_AMAX_MIN = 1e-4 # 1e-4 DSV4_SWA_PAGE_SIZE = 128 # 128 slots/page DSV4_C4_PAGE_SIZE = 64 # 64 slots/page DSV4_C128_PAGE_SIZE = 2 # 2 slots/page @@ -89,7 +89,7 @@ def write(self, layer_index: int, loc: torch.Tensor, packed: torch.Tensor) -> No if loc.numel() == 0: return loc = loc.reshape(-1) - packed = packed.reshape(-1, self.bytes_per_token).contiguous() + packed = packed.reshape(-1, self.bytes_per_token) flat = self.buffer[layer_index].view(-1) data_offsets, scale_offsets = self._loc_offsets(loc) data_range = torch.arange(self.data_bytes_per_token, device=loc.device) @@ -108,7 +108,7 @@ def read(self, layer_index: int, loc: torch.Tensor) -> torch.Tensor: scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device) data = flat[data_offsets.unsqueeze(1) + data_range.unsqueeze(0)] scale = flat[scale_offsets.unsqueeze(1) + scale_range.unsqueeze(0)] - return torch.cat([data, scale], dim=1).contiguous() + return torch.cat([data, scale], dim=1) class DeepseekV4MemoryManager(MemoryManager): @@ -591,11 +591,13 @@ def free_c128(self, free_index) -> None: # ------------------------------------------------------------------ packed codecs (torch reference) # 与 sglang/vllm 的 fp8_ds_mla 字节布局逐位对齐(ue8m0 幂次 scale)。这些 torch 实现是该 ABI 的 # 可执行规格(单测 oracle,triton writer 与其逐字节对拍),不可删除。 + # torch参考,未使用 def _pack_mla_kv(self, kv: torch.Tensor) -> torch.Tensor: kv = kv.reshape(-1, self.mla_head_dim) out = torch.empty((kv.shape[0], self.mla_bytes_per_token), dtype=torch.uint8, device=kv.device) nope = kv[:, : self.mla_nope_dim].float().reshape(-1, self.mla_scale_bytes - 1, self.mla_quant_group_size) - scale = torch.clamp(nope.abs().amax(dim=-1) / DSV4_FP8_E4M3_MAX, min=DSV4_FP8_SCALE_MIN) + amax = torch.clamp(nope.abs().amax(dim=-1), min=DSV4_FP8_AMAX_MIN) + scale = amax / DSV4_FP8_E4M3_MAX scale_exp = torch.ceil(torch.log2(scale)).to(torch.int32) scale = torch.exp2(scale_exp.float()) nope_fp8 = torch.clamp(nope / scale.unsqueeze(-1), -DSV4_FP8_E4M3_MAX, DSV4_FP8_E4M3_MAX).to( @@ -604,7 +606,7 @@ def _pack_mla_kv(self, kv: torch.Tensor) -> torch.Tensor: out[:, : self.mla_nope_dim].copy_(nope_fp8.reshape(-1, self.mla_nope_dim).view(dtype=torch.uint8)) rope_start = self.mla_nope_dim rope_end = rope_start + self.mla_rope_dim * 2 - rope = kv[:, self.mla_nope_dim : self.mla_head_dim].contiguous().to(torch.bfloat16) + rope = kv[:, self.mla_nope_dim : self.mla_head_dim].to(torch.bfloat16) out[:, rope_start:rope_end].copy_(rope.view(dtype=torch.uint8).reshape(-1, self.mla_rope_dim * 2)) scale_start = rope_end scale_end = scale_start + self.mla_scale_bytes - 1 @@ -636,10 +638,8 @@ def _pack_indexer_k(self, indexer_k: torch.Tensor) -> torch.Tensor: device=indexer_k.device, ) k_float = indexer_k.float() - scale = torch.clamp( - k_float.abs().amax(dim=-1, keepdim=True) / DSV4_FP8_E4M3_MAX, - min=DSV4_FP8_SCALE_MIN, - ) + amax = torch.clamp(k_float.abs().amax(dim=-1, keepdim=True), min=DSV4_FP8_AMAX_MIN) + scale = amax / DSV4_FP8_E4M3_MAX k_fp8 = torch.clamp(k_float / scale, -DSV4_FP8_E4M3_MAX, DSV4_FP8_E4M3_MAX).to(torch.float8_e4m3fn) out[:, : self.indexer_head_dim].copy_(k_fp8.view(dtype=torch.uint8)) out[:, self.indexer_head_dim :].copy_(scale.view(dtype=torch.uint8).reshape(-1, DSV4_INDEXER_SCALE_BYTES)) @@ -684,13 +684,11 @@ def pack_mla_kv_to_cache_fused_norm_rope( ): """同 pack_mla_kv_to_cache,但 rmsnorm + 尾部交错 rope 融合进写入 kernel 并省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" - if kv.shape[0] == 0: - return from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import ( fused_k_norm_rope_flashmla, ) - swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] + swa_slots = self.full_to_swa_indexs[mem_index.reshape(-1)] swa_slots = torch.where(swa_slots < 0, torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), swa_slots) fused_k_norm_rope_flashmla( kv=kv, @@ -719,7 +717,13 @@ def pack_compressed_kv_to_cache(self, layer_index: int, slots: torch.Tensor, com pool.page_size, ) - def pack_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer_k: torch.Tensor): + def pack_indexer_k_to_cache( + self, + layer_index: int, + mem_index: torch.Tensor, + positions: torch.Tensor, + indexer_k: torch.Tensor, + ): if indexer_k.shape[0] == 0: return assert self.compress_rates[layer_index] == 4, "只有 c4(CSA) 层有 indexer-K" @@ -729,7 +733,9 @@ def pack_indexer_k_to_cache(self, layer_index: int, slots: torch.Tensor, indexer destindex_copy_indexer_k_dsv4( indexer_k.reshape(-1, self.indexer_head_dim), - slots.to(indexer_k.device), + mem_index.reshape(-1), + positions.reshape(-1), + self.full_to_c4_indexs, self.c4_indexer_pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]), self.c4_indexer_pool.page_size, ) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index a270a8ac5f..2fd3228331 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -125,14 +125,14 @@ def _fused_compress_norm_rope_insert_kernel( STATE_RING: tl.constexpr, ROPE_HEAD_DIM: tl.constexpr, FP8_MAX: tl.constexpr, - SCALE_MIN: tl.constexpr, + AMAX_MIN: tl.constexpr, NOPE_DIM: tl.constexpr, QUANT_BLOCK: tl.constexpr, SCALE_BYTES: tl.constexpr, PAGE_SIZE: tl.constexpr, BYTES_PER_PAGE: tl.constexpr, BLOCK: tl.constexpr, - OUTPUT_BF16: tl.constexpr, + IS_IN_INDEXER: tl.constexpr, ): token_idx = tl.program_id(0) out_slot = tl.load(out_slots + token_idx).to(tl.int64) @@ -155,25 +155,33 @@ def _fused_compress_norm_rope_insert_kernel( q_start = token_idx - mtp_index token_offsets = tl.arange(0, WINDOW_SIZE) + # position 11, start 4 start = position - WINDOW_SIZE + 1 + # 最近的window个索引 [4, 5, 6, 7, 8, 9, 10, 11] gather_pos = start + token_offsets + # 都在[0, 12)之间,valid [1, 1, 1, 1, 1, 1, 1, 1] valid_pos = (gather_pos >= 0) & (gather_pos < seq_len) + # 新计算的mask ready_len=8 [0, 0, 0, 0, 1, 1, 1, 1] use_current = (gather_pos >= ready_len) & valid_pos + # 现在的kvscore的索引 [-4, -3, -2, -1, 0, 1, 2, 3] current_idx = q_start + (gather_pos - ready_len) + # cache的mask [1, 1, 1, 1, 0, 0, 0, 0] + valid_pos = valid_pos & (~use_current) + cache_pos = valid_pos if IS_C4: full_slot = tl.load( req_to_token + req_idx * req_to_token_stride0 + gather_pos, - mask=valid_pos & (~use_current), + mask=cache_pos, other=0, ).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot, mask=valid_pos & (~use_current), other=-1).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot, mask=cache_pos, other=-1).to(tl.int64) state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) - state_valid = valid_pos & (~use_current) & (swa_slot >= 0) + state_valid = cache_pos & (swa_slot >= 0) head_offset = tl.where(token_offsets >= COMPRESS_RATIO, HEAD_DIM, 0) else: state_row = req_idx * STATE_RING + gather_pos % STATE_RING - state_valid = valid_pos & (~use_current) + state_valid = cache_pos head_offset = token_offsets * 0 offs = tl.arange(0, BLOCK) @@ -181,25 +189,28 @@ def _fused_compress_norm_rope_insert_kernel( current_mask = use_current[:, None] & dim_mask[None, :] state_mask = state_valid[:, None] & dim_mask[None, :] + cur_ptrs = ( + kv_score + current_idx[:, None] * kv_score_stride0 + (head_offset[:, None] + offs[None, :]) * kv_score_stride1 + ) cur_kv = tl.load( - kv_score + current_idx[:, None] * kv_score_stride0 + (head_offset[:, None] + offs[None, :]) * kv_score_stride1, + cur_ptrs, mask=current_mask, other=0.0, ) cur_score = tl.load( - kv_score - + current_idx[:, None] * kv_score_stride0 - + (STATE_WIDTH + head_offset[:, None] + offs[None, :]) * kv_score_stride1, + cur_ptrs + STATE_WIDTH * kv_score_stride1, mask=current_mask, other=float("-inf"), ) + + state_ptrs = state_buffer + state_row[:, None] * STATE_LAST_DIM + head_offset[:, None] + offs[None, :] state_kv = tl.load( - state_buffer + state_row[:, None] * STATE_LAST_DIM + head_offset[:, None] + offs[None, :], + state_ptrs, mask=state_mask, other=0.0, ) state_score = tl.load( - state_buffer + state_row[:, None] * STATE_LAST_DIM + STATE_WIDTH + head_offset[:, None] + offs[None, :], + state_ptrs + STATE_WIDTH, mask=state_mask, other=float("-inf"), ) @@ -208,12 +219,15 @@ def _fused_compress_norm_rope_insert_kernel( score = tl.where(current_mask, cur_score, state_score) score = tl.softmax(score, dim=0) compressed_kv = tl.sum(kv * score, axis=0) + # 以上得到compressed_kv + # rms_norm rms_w = tl.load(norm_weight + offs, mask=dim_mask, other=0.0) variance = tl.sum(compressed_kv * compressed_kv, axis=0) / HEAD_DIM rrms = tl.rsqrt(variance + rms_eps) normed = compressed_kv * rrms * rms_w + # rope num_pairs: tl.constexpr = BLOCK // 2 nope_pairs: tl.constexpr = NOPE_DIM // 2 pair_2d = tl.reshape(normed, (num_pairs, 2)) @@ -229,7 +243,7 @@ def _fused_compress_norm_rope_insert_kernel( new_odd = odd * cos_v + even * sin_v rotated = tl.interleave(new_even, new_odd) - if OUTPUT_BF16: + if IS_IN_INDEXER: # indexer-K path: emit the post-rope full HEAD_DIM vector as dense bf16 (token-indexed), # leaving the fp8 single-amax pack to destindex_copy_indexer_k_dsv4 (the c4_indexer_pool # ABI differs from the latent slab: whole-vector fp8 + one fp32 scale, no bf16 rope tail). @@ -243,16 +257,18 @@ def _fused_compress_norm_rope_insert_kernel( n_quant_blocks: tl.constexpr = BLOCK // QUANT_BLOCK n_nope_blocks: tl.constexpr = NOPE_DIM // QUANT_BLOCK + # 有待检验是否有必要 quant_input = normed.to(tl.bfloat16).to(tl.float32) quant_2d = tl.reshape(quant_input, (n_quant_blocks, QUANT_BLOCK)) abs_2d = tl.abs(quant_2d) - block_absmax = tl.max(abs_2d, axis=1) - scale_exp = tl.ceil(libdevice.log2(tl.maximum(block_absmax / FP8_MAX, SCALE_MIN))).to(tl.int32) + block_absmax = tl.maximum(tl.max(abs_2d, axis=1), AMAX_MIN) + scale_exp = tl.ceil(libdevice.log2(block_absmax / FP8_MAX)).to(tl.int32) scale = ((scale_exp + 127) << 23).to(tl.float32, bitcast=True) kv_fp8 = tl.clamp(quant_2d / scale[:, None], -FP8_MAX, FP8_MAX).to(tl.float8e4nv) kv_u8 = tl.reshape(kv_fp8.to(tl.uint8, bitcast=True), (BLOCK,)) tl.store(out_buffer + data_base + offs, kv_u8, mask=offs < NOPE_DIM) + # ue8m0 scale_idx = tl.arange(0, SCALE_BYTES) scale_bytes = tl.where(scale_idx < n_nope_blocks, scale_exp + 127, 0).to(tl.uint8) tl.store(out_buffer + scale_base + scale_idx, scale_bytes) @@ -264,15 +280,13 @@ def _fused_compress_norm_rope_insert_kernel( return -def prepare_partial_states( +def apply_ape( *, kv_score: torch.Tensor, position_ids: torch.Tensor, ape: torch.Tensor, compress_ratio: int, ): - if kv_score.shape[0] == 0: - return state_width = kv_score.shape[-1] // 2 _add_ape_to_kv_score_kernel[(kv_score.shape[0],)]( kv_score, @@ -301,10 +315,10 @@ def fused_compress( compress_ratio: int, cos_table: torch.Tensor, sin_table: torch.Tensor, - output_bf16: bool = False, + is_in_indexer: bool = False, ): mem_manager = infer_state.mem_manager - if output_bf16: + if is_in_indexer: assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) @@ -379,14 +393,14 @@ def fused_compress( STATE_RING=state_ring, ROPE_HEAD_DIM=qk_rope_head_dim, FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, - SCALE_MIN=1e-4, + AMAX_MIN=1e-4, NOPE_DIM=head_dim - qk_rope_head_dim, QUANT_BLOCK=64, SCALE_BYTES=(head_dim - qk_rope_head_dim) // 64 + 1, PAGE_SIZE=out_page_size, BYTES_PER_PAGE=out_buffer.shape[-1], BLOCK=block_head, - OUTPUT_BF16=output_bf16, + IS_IN_INDEXER=is_in_indexer, num_warps=4, ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 2bc482955f..e5a1fddba8 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -11,7 +11,7 @@ from lightllm.utils.vllm_utils import vllm_ops from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post from .compressor import fused_compress as fused_compress_op -from .compressor import prepare_partial_states +from .compressor import apply_ape from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from ..infer_struct import DeepseekV4InferStateInfo import deep_gemm @@ -456,7 +456,7 @@ def compress( norm_weight = layer_weight.compressor_norm_.weight ape = layer_weight.compressor_ape_.weight head_dim = self.head_dim - prepare_partial_states( + apply_ape( kv_score=kv_score, position_ids=infer_state.position_ids, ape=ape, @@ -473,7 +473,7 @@ def compress( compress_ratio=self.compress_ratio, cos_table=cos_table, sin_table=sin_table, - output_bf16=self.is_in_indexer, + is_in_indexer=self.is_in_indexer, ) @@ -517,10 +517,6 @@ def write_indexer_k( sin_table, use_custom_tensor_manager=True, ): - """c4-only: compress this step's tokens into per-c4-entry indexer keys and pack them into - c4_indexer_pool. MUST run before build_metadata so the scorer (gather + deep_gemm.fp8_mqa_logits) - reads the finished entries; runs every step (incl. in the decode graph) so keys accumulate for - later long-context scoring. No-op on c128 / dense layers.""" if self.compress_ratio != 4: return # Only group-end rows in this dense bf16 scratch are valid indexer keys. @@ -538,13 +534,12 @@ def write_indexer_k( scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5) mem_manager = infer_state.mem_manager - positions = infer_state.position_ids - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] - # only group-end tokens finish a c4 entry; mask the rest to -1 so the packer skips them - # (mid-group tokens share the group's c4 slot -> avoids racing a finished slot). - completed = ((positions + 1) % 4 == 0) & (out_slots >= 0) - masked_slots = torch.where(completed, out_slots, torch.full_like(out_slots, -1)).to(torch.int32) - mem_manager.pack_indexer_k_to_cache(self.layer_idx_, masked_slots, scratch) + mem_manager.pack_indexer_k_to_cache( + self.layer_idx_, + infer_state.mem_index.reshape(-1), + infer_state.position_ids, + scratch, + ) def build_metadata( self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True @@ -575,10 +570,6 @@ def build_metadata( def _indexer_q_weight( self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True ): - """fp8 indexer q (mirrors deepseek3_2 NsaInfer): wq_b -> rope(last rope dims) -> 1/sqrt(d) - Hadamard -> per-token fp8 quant. Returns (idx_q_fp8 [T,H,d], weights [T,H]); the per-token q - fp8 scale and the head_dim^-0.5 * n_heads^-0.5 score scale are folded into weights -- the - deep_gemm.fp8_mqa_logits contract (fp8 q carries no companion scale). Replicated -> full heads.""" # Fused: wq_b mm -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant, with the # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: # rotary_emb_fwd + hadamard_transform + act_quant + weights mul). freqs_cis is the compress rope diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 072591a552..3fc3dcf77d 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -331,7 +331,7 @@ def _dequant_in_place(self, weights): if woa in weights and weights[woa].dim() == 2: w = weights[woa] per_group_in = self.n_heads * self.head_dim // self.o_groups - weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2).contiguous() + weights[woa] = w.view(self.o_groups, self.o_lora_rank, per_group_in).transpose(1, 2) # Keep c4 overlap APE in checkpoint layout [4, 2*head_dim]. SGLang reorders it # because its compressor consumes ape.view(8, head_dim) by window offset. LightLLM # adds APE into each token's two score halves before compression using position % 4, diff --git a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py index 3510b92c30..25724000bc 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py @@ -7,21 +7,28 @@ @triton.jit def _fwd_kernel_destindex_copy_indexer_k_dsv4( K, - Dest_loc, + Mem_index, + Positions, + Full_to_c4, O_fp8, O_f32, stride_k_bs, stride_k_d, FP8_MIN: tl.constexpr, FP8_MAX: tl.constexpr, - SCALE_MIN: tl.constexpr, + AMAX_MIN: tl.constexpr, HEAD_DIM: tl.constexpr, + COMPRESS_RATIO: tl.constexpr, PAGE_SIZE: tl.constexpr, BYTES_PER_PAGE: tl.constexpr, ): cur_index = tl.program_id(0) - dest_index = tl.load(Dest_loc + cur_index).to(tl.int64) - # negative dest (unmapped slot) is a no-op, not an OOB write into a neighboring page. + position = tl.load(Positions + cur_index) + if (position + 1) % COMPRESS_RATIO != 0: + return + + full_slot = tl.load(Mem_index + cur_index).to(tl.int64) + dest_index = tl.load(Full_to_c4 + full_slot).to(tl.int64) if dest_index < 0: return @@ -30,9 +37,9 @@ def _fwd_kernel_destindex_copy_indexer_k_dsv4( offs_d = tl.arange(0, HEAD_DIM) vals = tl.load(K + cur_index * stride_k_bs + offs_d * stride_k_d).to(tl.float32) - amax = tl.max(tl.abs(vals), axis=0) + amax = tl.maximum(tl.max(tl.abs(vals), axis=0), AMAX_MIN) # per-token plain fp32 scale (not ue8m0), matching DeepseekV4MemoryManager._pack_indexer_k - scale = tl.maximum(amax / FP8_MAX, SCALE_MIN) + scale = amax / FP8_MAX k_fp8 = tl.clamp(vals / scale, min=FP8_MIN, max=FP8_MAX).to(tl.float8e4nv) data_base = page * BYTES_PER_PAGE + token_in_page * HEAD_DIM @@ -45,27 +52,32 @@ def _fwd_kernel_destindex_copy_indexer_k_dsv4( @torch.no_grad() def destindex_copy_indexer_k_dsv4( K: torch.Tensor, - DestLoc: torch.Tensor, + MemIndex: torch.Tensor, + Positions: torch.Tensor, + FullToC4: torch.Tensor, O_buffer: torch.Tensor, page_size: int, ): """Packed indexer-K page-slab writer (DeepSeek-V4 c4/CSA layers). K: [T, 128] bf16 unquantized indexer keys. - DestLoc: [T] int — c4-pool-local token slots; must already be allocated by the caller. - Negative slots (unmapped) are skipped. + MemIndex: [T] int — full-token slots for the current rows. + Positions: [T] int — logical token positions; only c4 group-end rows are written. + FullToC4: [full_pool_size + 1] int — full-token slot to c4-pool slot mapping. + Negative mappings are skipped. O_buffer: [num_pages, bytes_per_page] uint8 — one layer's slab from the c4 indexer PackedPagePool (128B fp8 data region + 4B fp32 scale tail per token). Bit-compatible with DeepseekV4MemoryManager._pack_indexer_k + PackedPagePool.write. """ - seq_len = DestLoc.shape[0] + seq_len = MemIndex.shape[0] if seq_len == 0: return head_dim, scale_bytes = 128, 4 K = K.reshape(-1, head_dim) assert K.shape[0] == seq_len, f"Expected K shape[0]={seq_len}, got {K.shape[0]}" + assert Positions.numel() == seq_len, f"Expected {seq_len} positions, got {Positions.numel()}" assert K.dtype == torch.bfloat16, f"Expected bf16 indexer K, got {K.dtype}" bytes_per_page = O_buffer.shape[-1] assert O_buffer.dtype == torch.uint8 and O_buffer.is_contiguous() @@ -75,15 +87,18 @@ def destindex_copy_indexer_k_dsv4( flat = O_buffer.view(-1) _fwd_kernel_destindex_copy_indexer_k_dsv4[(seq_len,)]( K, - DestLoc, + MemIndex, + Positions, + FullToC4, flat.view(torch.float8_e4m3fn), flat.view(torch.float32), K.stride(0), K.stride(1), FP8_MIN=torch.finfo(torch.float8_e4m3fn).min, FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, - SCALE_MIN=1e-4, + AMAX_MIN=1e-4, HEAD_DIM=head_dim, + COMPRESS_RATIO=4, PAGE_SIZE=page_size, BYTES_PER_PAGE=bytes_per_page, num_warps=1, diff --git a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py index a3ec6ed8cf..054d4a3ca2 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_kv_flashmla_dsv4.py @@ -16,7 +16,7 @@ def _fwd_kernel_destindex_copy_kv_flashmla_dsv4( stride_kv_d, FP8_MIN: tl.constexpr, FP8_MAX: tl.constexpr, - SCALE_MIN: tl.constexpr, + AMAX_MIN: tl.constexpr, NOPE_DIM: tl.constexpr, ROPE_DIM: tl.constexpr, GROUP_SIZE: tl.constexpr, @@ -46,8 +46,8 @@ def _fwd_kernel_destindex_copy_kv_flashmla_dsv4( group_mask = offs_g < NUM_GROUPS kv_ptrs = KV + cur_index * stride_kv_bs + (offs_g[:, None] * GROUP_SIZE + offs_e[None, :]) * stride_kv_d vals = tl.load(kv_ptrs, mask=group_mask[:, None], other=0.0).to(tl.float32) - amax = tl.max(tl.abs(vals), axis=1) - scale_exp = tl.ceil(libdevice.log2(tl.maximum(amax / FP8_MAX, SCALE_MIN))).to(tl.int32) + amax = tl.maximum(tl.max(tl.abs(vals), axis=1), AMAX_MIN) + scale_exp = tl.ceil(libdevice.log2(amax / FP8_MAX)).to(tl.int32) scale = ((scale_exp + 127) << 23).to(tl.float32, bitcast=True) kv_fp8 = tl.clamp(vals / scale[:, None], min=FP8_MIN, max=FP8_MAX).to(tl.float8e4nv) tl.store(O_fp8 + data_base + offs_g[:, None] * GROUP_SIZE + offs_e[None, :], kv_fp8, mask=group_mask[:, None]) @@ -107,7 +107,7 @@ def destindex_copy_kv_flashmla_dsv4( KV.stride(1), FP8_MIN=torch.finfo(torch.float8_e4m3fn).min, FP8_MAX=torch.finfo(torch.float8_e4m3fn).max, - SCALE_MIN=1e-4, + AMAX_MIN=1e-4, NOPE_DIM=nope_dim, ROPE_DIM=rope_dim, GROUP_SIZE=group_size, From a45e07d8939f67513f9db7c16b817c7e5711f15f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 23 Jul 2026 13:57:44 +0000 Subject: [PATCH 094/214] refact metadata --- .../triton_kernel/hadamard_transform.py | 9 +- .../deepseek_v4/layer_infer/compressor.py | 12 +- .../layer_infer/transformer_layer_infer.py | 128 +++++++++++------- .../triton_kernel/norm_rope_cuda.py | 8 +- 4 files changed, 94 insertions(+), 63 deletions(-) diff --git a/lightllm/models/deepseek3_2/triton_kernel/hadamard_transform.py b/lightllm/models/deepseek3_2/triton_kernel/hadamard_transform.py index eabf703f56..4784e8197e 100644 --- a/lightllm/models/deepseek3_2/triton_kernel/hadamard_transform.py +++ b/lightllm/models/deepseek3_2/triton_kernel/hadamard_transform.py @@ -51,13 +51,14 @@ def _pick_block_r(rows: int, device_index: int) -> int: return max(1, min(128, block_r)) -def _hadamard_transform_triton(x: torch.Tensor, scale: float) -> torch.Tensor: +def _hadamard_transform_triton(x: torch.Tensor, scale: float, out: torch.Tensor = None) -> torch.Tensor: original_shape = x.shape hidden_size = x.size(-1) if not x.is_contiguous(): x = x.contiguous() rows = x.numel() // hidden_size - out = torch.empty_like(x) + if out is None: + out = torch.empty_like(x) BLOCK_R = _pick_block_r(rows, x.device.index) grid = (triton.cdiv(rows, BLOCK_R),) _hadamard_transform_kernel[grid]( @@ -72,9 +73,9 @@ def _hadamard_transform_triton(x: torch.Tensor, scale: float) -> torch.Tensor: return out.view(original_shape) -def hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor: +def hadamard_transform(x: torch.Tensor, scale: float = 1.0, out: torch.Tensor = None) -> torch.Tensor: assert x.is_cuda, "hadamard_transform only supports CUDA tensors" assert x.dtype == torch.bfloat16, "hadamard_transform expects bfloat16 input" assert x.size(-1) == 128, "DeepSeek-V3.2 Hadamard transform expects hidden size 128" - return _hadamard_transform_triton(x, scale) + return _hadamard_transform_triton(x, scale, out) diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 2fd3228331..5b675ec882 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -316,6 +316,7 @@ def fused_compress( cos_table: torch.Tensor, sin_table: torch.Tensor, is_in_indexer: bool = False, + out_buffer: torch.Tensor = None, ): mem_manager = infer_state.mem_manager if is_in_indexer: @@ -323,11 +324,12 @@ def fused_compress( out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) state_ring = mem_manager.c4_state_ring - out_buffer = torch.empty( - (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), - dtype=torch.bfloat16, - device=infer_state.mem_index.device, - ) + if out_buffer is None: + out_buffer = torch.empty( + (infer_state.mem_index.numel(), mem_manager.indexer_head_dim), + dtype=torch.bfloat16, + device=infer_state.mem_index.device, + ) out_page_size = 1 else: if compress_ratio == 4: diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index e5a1fddba8..9b75bfe78d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,6 +1,6 @@ import torch import torch.distributed as dist -from lightllm.common.basemodel import TransformerLayerInferTpl +from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.distributed.communication_op import all_reduce @@ -412,7 +412,7 @@ def _select_experts( return weights, indices -class CompressorInfer: +class CompressorInfer(BaseLayerInfer): """Window-softmax compressor. is_in_indexer=False compresses the c4/c128 latent KV into the paged fp8 slab (attention extra_k); is_in_indexer=True reuses the SAME machinery (mirroring sglang's Compressor(is_in_indexer=...)) with the indexer weights/dims/state pool to produce the @@ -462,6 +462,13 @@ def compress( ape=ape, compress_ratio=self.compress_ratio, ) + out_buffer = None + if self.is_in_indexer and use_custom_tensor_manager: + out_buffer = self.alloc_tensor( + (infer_state.mem_index.numel(), self.index_head_dim), + torch.bfloat16, + device=infer_state.mem_index.device, + ) return fused_compress_op( kv_score=kv_score, infer_state=infer_state, @@ -474,10 +481,11 @@ def compress( cos_table=cos_table, sin_table=sin_table, is_in_indexer=self.is_in_indexer, + out_buffer=out_buffer, ) -class DeepseekV4IndexInfer: +class DeepseekV4IndexInfer(BaseLayerInfer): """Model-side builder for the FlashMLA sparse-index metadata. Mirrors deepseek3_2's NsaInfer boundary (the model owns ALL index construction; the attention backend only forwards final tensors to flash_mla.flash_mla_with_kvcache) AND its c4 implementation: hadamard'd fp8 q/K, a @@ -532,7 +540,10 @@ def write_indexer_k( # (Hq)·(Hk)=q·k (H orthogonal) and the fp8 quant of K stays accurate. from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform - scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5) + hadamard_out = None + if use_custom_tensor_manager: + hadamard_out = self.alloc_tensor(scratch.shape, scratch.dtype, device=scratch.device) + scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5, out=hadamard_out) mem_manager = infer_state.mem_manager mem_manager.pack_indexer_k_to_cache( self.layer_idx_, @@ -544,10 +555,6 @@ def write_indexer_k( def build_metadata( self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True ): - """Return the final flash_mla index tensors for this layer's compress variant. swa indices and - the per-token req_idx are layer-independent and precomputed once in init_some_extra_state - (read here); only the c4 scorer is per-layer. The backend pairs these with the - (data-independent, layer-keyed) fp8 cache-byte views it owns.""" swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) swa_lengths = infer_state.dsv4_swa_lengths positions = infer_state.position_ids @@ -589,12 +596,18 @@ def _indexer_q_weight( raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( token_num, self.index_n_heads ) # [T, H] raw + idx_q_fp8_out = weights_out = None + if use_custom_tensor_manager: + idx_q_fp8_out = self.alloc_tensor(idx_q.shape, torch.float8_e4m3fn, device=idx_q.device) + weights_out = self.alloc_tensor((*idx_q.shape[:-1], 1), torch.float32, device=idx_q.device) idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( idx_q, raw_w, self.indexer_weight_scale, self.freqs_cis, infer_state.position_ids, + q_fp8=idx_q_fp8_out, + weights_out=weights_out, ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded return idx_q_fp8, weights.squeeze(-1) @@ -608,9 +621,6 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, max_entries = max(1, int(infer_state.max_kv_seq_len) // 4) c4_cap = ((max_entries + 63) // 64) * 64 - # entry space fits the budget -> every causal entry is selected; no scoring needed. The - # captured decode graph (graph_max_len -> max_entries > topk) always takes the scorer branch - # below, so this only shortcuts tiny eager contexts. if max_entries <= index_topk: from ..triton_kernel.build_compress_index_dsv4 import build_compress_index @@ -626,23 +636,16 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, ) return slots.unsqueeze(1), lengths - c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req - device = positions.device page_size = mem_manager.c4_indexer_pool.page_size - # The page table / row_page_table / valid_len / ctx_lens / paged-logits metadata / topk_lengths - # are LAYER-INDEPENDENT (depend on request layout + c4_cap, not on weights/layer). Build them on - # the first c4 layer of the forward and reuse on the other ~20 c4 layers (was rebuilt per layer: - # build_c4_indexer_page_table + a [T,npages] gather + clamp/reshape + get_paged_mqa_logits_metadata - # each, i.e. ~20x redundant index/copy/clamp launches). Lazy (not init_some_extra_state) so it is - # computed inside the decode cuda graph with the capture-forced shapes -> no graph-cap mismatch. cached = getattr(infer_state, "_c4_paged_meta", None) if cached is None: from ..triton_kernel.gather_c4_indexer_k_dsv4 import build_c4_indexer_page_table b_req_idx = infer_state.b_req_idx batch = b_req_idx.shape[0] + c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req page_table = build_c4_indexer_page_table( mem_manager, b_req_idx, @@ -664,50 +667,71 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, valid_len = ((positions + 1) // 4).to(torch.int32) ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1) - metadata = deep_gemm.get_paged_mqa_logits_metadata( - ctx_lens, - page_size, - deep_gemm.get_num_sms(), + rows_per_chunk = None + chunk_metadata = None + if infer_state.is_prefill: + rows_per_chunk = max(1, _C4_PREFILL_LOGITS_BUDGET_BYTES // (c4_cap * 4)) + if positions.numel() > rows_per_chunk: + chunk_metadata = tuple( + deep_gemm.get_paged_mqa_logits_metadata( + ctx_lens[start : start + rows_per_chunk], + page_size, + deep_gemm.get_num_sms(), + ) + for start in range(0, positions.numel(), rows_per_chunk) + ) + + metadata = ( + deep_gemm.get_paged_mqa_logits_metadata( + ctx_lens, + page_size, + deep_gemm.get_num_sms(), + ) + if chunk_metadata is None + else None ) topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) - cached = (row_page_table, valid_len, ctx_lens, metadata, topk_lengths) + cached = ( + row_page_table, + valid_len, + ctx_lens, + metadata, + topk_lengths, + rows_per_chunk, + chunk_metadata, + ) infer_state._c4_paged_meta = cached - row_page_table, valid_len, ctx_lens, metadata, topk_lengths = cached - kv_cache = mem_manager.c4_indexer_pool.get_layer_buffer(mem_manager.layer_to_c4_idx[self.layer_idx_]).view( + row_page_table, valid_len, ctx_lens, metadata, topk_lengths, rows_per_chunk, chunk_metadata = cached + indexer_k_cache = mem_manager.c4_indexer_pool.get_layer_buffer( + mem_manager.layer_to_c4_idx[self.layer_idx_] + ).view( mem_manager.c4_indexer_pool.num_pages, page_size, 1, self.index_head_dim + 4, ) top_slots, _ = workspace.c4(infer_state.microbatch_index, idx_q_fp8.shape[0], index_topk) - if infer_state.is_prefill: - rows_per_chunk = max(1, _C4_PREFILL_LOGITS_BUDGET_BYTES // (c4_cap * 4)) - if idx_q_fp8.shape[0] > rows_per_chunk: - for start in range(0, idx_q_fp8.shape[0], rows_per_chunk): - end = min(start + rows_per_chunk, idx_q_fp8.shape[0]) - chunk_ctx_lens = ctx_lens[start:end] - self._c4_score_topk( - idx_q_fp8[start:end], - kv_cache, - weights[start:end], - chunk_ctx_lens, - row_page_table[start:end], - deep_gemm.get_paged_mqa_logits_metadata( - chunk_ctx_lens, - page_size, - deep_gemm.get_num_sms(), - ), - c4_cap, - valid_len[start:end], - top_slots[start:end], - page_size, - ) - return top_slots.unsqueeze(1), topk_lengths + if chunk_metadata is not None: + for chunk_idx, start in enumerate(range(0, idx_q_fp8.shape[0], rows_per_chunk)): + end = min(start + rows_per_chunk, idx_q_fp8.shape[0]) + self._c4_score_topk( + idx_q_fp8[start:end], + indexer_k_cache, + weights[start:end], + ctx_lens[start:end], + row_page_table[start:end], + chunk_metadata[chunk_idx], + c4_cap, + valid_len[start:end], + top_slots[start:end], + page_size, + ) + return top_slots.unsqueeze(1), topk_lengths self._c4_score_topk( idx_q_fp8, - kv_cache, + indexer_k_cache, weights, ctx_lens, row_page_table, @@ -722,7 +746,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, @staticmethod def _c4_score_topk( idx_q_fp8, - kv_cache, + indexer_k_cache, weights, ctx_lens, row_page_table, @@ -734,7 +758,7 @@ def _c4_score_topk( ): logits = deep_gemm.fp8_paged_mqa_logits( idx_q_fp8.unsqueeze(1), - kv_cache, + indexer_k_cache, weights, ctx_lens, row_page_table, diff --git a/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py b/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py index 0ef4434f2e..3984a398ef 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py +++ b/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py @@ -83,9 +83,13 @@ def fused_q_indexer_rope_hadamard_quant( weight_scale: float, freqs_cis: torch.Tensor, positions: torch.Tensor, + q_fp8: torch.Tensor = None, + weights_out: torch.Tensor = None, ): - q_fp8 = torch.empty(q_input.shape, dtype=torch.float8_e4m3fn, device=q_input.device) - weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) + if q_fp8 is None: + q_fp8 = torch.empty(q_input.shape, dtype=torch.float8_e4m3fn, device=q_input.device) + if weights_out is None: + weights_out = torch.empty((*q_input.shape[:-1], 1), dtype=torch.float32, device=q_input.device) _load_cuda().fused_q_indexer_rope_hadamard_quant( q_input, q_fp8, From 82c2bfeeaba9cee7d3ea58b51affc438d0bfdab5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 23 Jul 2026 14:13:33 +0000 Subject: [PATCH 095/214] ep delete shared reduced --- .../layer_infer/transformer_layer_infer.py | 12 +----------- .../layer_weights/transformer_layer_weight.py | 9 +++++++++ 2 files changed, 10 insertions(+), 11 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 9b75bfe78d..3cc2cc1887 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,9 +1,7 @@ import torch -import torch.distributed as dist from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd -from lightllm.distributed.communication_op import all_reduce from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args @@ -357,8 +355,7 @@ def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: De def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): x = x.view(-1, self.embed_dim_) - if not self.enable_ep_moe: - x = self._tpsp_allgather(input=x, infer_state=infer_state) + x = self._tpsp_allgather(input=x, infer_state=infer_state) logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) weights, indices = self._select_experts(logits, infer_state, layer_weight) @@ -369,13 +366,6 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) if self.enable_ep_moe: - if self.tp_world_size_ > 1: - all_reduce( - shared, - op=dist.ReduceOp.SUM, - group=infer_state.dist_group, - async_op=False, - ) return routed + shared out = routed + shared return self._tpsp_reduce(input=out, infer_state=infer_state) diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 3fc3dcf77d..493fb1cf27 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -1,5 +1,6 @@ import torch from lightllm.common.basemodel import TransformerLayerWeight +from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.basemodel.layer_weights.meta_weights import ( ROWMMWeight, COLMMWeight, @@ -217,12 +218,18 @@ def _init_moe(self): # silu_and_mul triton kernel, no swiglu clamp) drives it directly. Order [w1, w3] = [gate, up] # matches silu_and_mul_fwd's blocked layout (first half gate, second half up). sp = f"{p}.shared_experts" + # EP keeps a complete shared expert per rank, so its output needs no per-layer all-reduce. + enable_ep_moe = get_env_start_args().enable_ep_moe + shared_tp_rank = 0 if enable_ep_moe else None + shared_tp_world_size = 1 if enable_ep_moe else None self.gate_up_proj = ROWMMWeight( in_dim=self.hidden, out_dims=[self.moe_inter, self.moe_inter], weight_names=[f"{sp}.w1.weight", f"{sp}.w3.weight"], data_type=self.data_type_, quant_method=self.get_quant_method("shared_gate"), + tp_rank=shared_tp_rank, + tp_world_size=shared_tp_world_size, ) self.down_proj = COLMMWeight( in_dim=self.moe_inter, @@ -230,6 +237,8 @@ def _init_moe(self): weight_names=f"{sp}.w2.weight", data_type=self.data_type_, quant_method=self.get_quant_method("shared_down"), + tp_rank=shared_tp_rank, + tp_world_size=shared_tp_world_size, ) self.experts_ = FusedMoeWeight( gate_proj_name="w1", From e726a699d7efaf1c587db0f98e4ac40437bb86ec Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 3 Aug 2026 15:47:02 +0000 Subject: [PATCH 096/214] support cpucache --- .../deepseek4_mem_manager.py | 357 +++++++- .../common/kv_cache_mem_manager/mem_utils.py | 9 + .../kv_cache_mem_manager/operator/deepseek.py | 38 + lightllm/common/req_manager.py | 5 + lightllm/models/deepseek_v4/model.py | 6 + .../deepseek_v4/triton_kernel/cpu_cache_io.py | 765 ++++++++++++++++++ lightllm/server/api_cli.py | 4 +- lightllm/server/api_start.py | 8 + lightllm/server/core/objs/start_args_type.py | 2 +- .../multi_level_kv_cache/cpu_cache_client.py | 45 +- .../model_infer/mode_backend/base_backend.py | 18 +- .../mode_backend/dsv4_multi_level_kv_cache.py | 386 +++++++++ .../mode_backend/multi_level_kv_cache.py | 8 +- lightllm/utils/envs_utils.py | 5 + lightllm/utils/kv_cache_utils.py | 47 +- 15 files changed, 1666 insertions(+), 37 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py create mode 100644 lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index a5092cb08d..e56ec095a0 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,5 +1,6 @@ import torch -from typing import List, Optional, Union +from dataclasses import dataclass +from typing import List, Optional, Sequence, Union from .mem_manager import MemoryManager from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator @@ -12,7 +13,7 @@ # fp8_ds_mla packed-latent byte layout (ABI shared with the flash_mla extra-cache fork and # sglang/vllm): 448B NoPE fp8 + 64*2B RoPE bf16 + 7B ue8m0 scale + 1B pad = 584B per token, -# stored in page slabs whose tail carries the per-token scale bytes. +# stored in packed GPU pages whose tail carries the per-token scale bytes. DSV4_MLA_NOPE_DIM = 448 # 448B DSV4_MLA_ROPE_DIM = 64 # 64 dim DSV4_MLA_HEAD_DIM = DSV4_MLA_NOPE_DIM + DSV4_MLA_ROPE_DIM # 512 @@ -30,6 +31,7 @@ DSV4_C4_PAGE_SIZE = 64 # 64 slots/page DSV4_C128_PAGE_SIZE = 2 # 2 slots/page DSV4_PROMPT_CACHE_PAGE_SIZE = DSV4_C4_PAGE_SIZE * 4 # 256 (= c4 ratio) +DSV4_CPU_CACHE_TOKEN_PAGE_SIZE = 2048 # compressor state ring: c4 overlap 对的基础窗口为每页 2 个分组槽 × ratio 4 行;MTP # 追加候选槽,避免 rejected draft 覆盖仍存活的基础窗口。c128 同样追加候选槽,再对齐到 ratio 4。 DSV4_C4_STATE_RING = 8 # 8 rows/page before MTP padding @@ -43,8 +45,196 @@ def _ceil_div(a: int, b: int) -> int: return (a + b - 1) // b +def _aligned_gpu_page_nbytes(page_size: int, data_nbytes: int, scale_nbytes: int, align_nbytes: int = 1) -> int: + return _ceil_div(page_size * (data_nbytes + scale_nbytes), align_nbytes) * align_nbytes + + +@dataclass(frozen=True) +class DeepseekV4CpuCacheLayout: + """Opaque CPU checkpoint ABI for the composite DeepSeek-V4 cache. + + Every checkpoint contains compressed history for ``token_page_size`` tokens + and only the final 256 tokens of SWA/continuation state. The history is + arranged in independent 256-token blocks so a radix hit can skip an initial + part of a checkpoint without copying or allocating it again. + """ + + token_page_size: int + history_block_num: int + layer_num: int + n_c4: int + n_c128: int + head_dim: int + indexer_head_dim: int + + c4_offset: int + c4_gpu_page_nbytes: int + c4_gpu_pages_per_cpu_page: int + c4_layer_nbytes: int + c4_nbytes: int + + c4_indexer_offset: int + c4_indexer_gpu_page_nbytes: int + c4_indexer_layer_nbytes: int + c4_indexer_nbytes: int + + c128_offset: int + c128_row_nbytes: int + c128_rows_per_page: int + c128_layer_nbytes: int + c128_nbytes: int + + swa_offset: int + swa_gpu_page_nbytes: int + swa_gpu_pages_per_cpu_page: int + swa_layer_nbytes: int + swa_nbytes: int + + c4_state_offset: int + c4_state_row_nbytes: int + c4_state_rows: int + c4_state_nbytes: int + + c4_indexer_state_offset: int + c4_indexer_state_row_nbytes: int + c4_indexer_state_nbytes: int + page_nbytes: int + + @classmethod + def from_compress_rates( + cls, + compress_rates: Sequence[int], + token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, + head_dim: int = DSV4_MLA_HEAD_DIM, + indexer_head_dim: int = DSV4_INDEXER_HEAD_DIM, + ) -> "DeepseekV4CpuCacheLayout": + rates = tuple(int(rate) for rate in compress_rates) + if token_page_size <= 0 or token_page_size % DSV4_PROMPT_CACHE_PAGE_SIZE != 0: + raise ValueError( + f"DeepSeek-V4 CPU cache token page size must be a positive multiple of " + f"{DSV4_PROMPT_CACHE_PAGE_SIZE}, got {token_page_size}" + ) + if head_dim != DSV4_MLA_HEAD_DIM: + raise ValueError(f"DeepSeek-V4 CPU cache expects head_dim={DSV4_MLA_HEAD_DIM}, got {head_dim}") + if indexer_head_dim != DSV4_INDEXER_HEAD_DIM: + raise ValueError( + f"DeepSeek-V4 CPU cache expects indexer_head_dim={DSV4_INDEXER_HEAD_DIM}, got {indexer_head_dim}" + ) + + layer_num = len(rates) + n_c4 = rates.count(4) + n_c128 = rates.count(128) + history_block_num = token_page_size // DSV4_PROMPT_CACHE_PAGE_SIZE + + # 64 * (576B data + 8B scale) = 37,376B; align to 576B -> 37,440B. + c4_gpu_page_nbytes = _aligned_gpu_page_nbytes( + DSV4_C4_PAGE_SIZE, + DSV4_MLA_DATA_BYTES_PER_TOKEN, + DSV4_MLA_SCALE_BYTES, + DSV4_MLA_PAGE_ALIGN_BYTES, + ) + c4_gpu_pages_per_cpu_page = history_block_num + c4_layer_nbytes = c4_gpu_pages_per_cpu_page * c4_gpu_page_nbytes + c4_offset = 0 + c4_nbytes = n_c4 * c4_layer_nbytes + + # indexer_head_dim is validated as 128: 64 * (128B data + 4B scale) = 8,448B. + c4_indexer_gpu_page_nbytes = _aligned_gpu_page_nbytes( + DSV4_C4_PAGE_SIZE, + indexer_head_dim, + DSV4_INDEXER_SCALE_BYTES, + ) + c4_indexer_layer_nbytes = c4_gpu_pages_per_cpu_page * c4_indexer_gpu_page_nbytes + c4_indexer_offset = c4_offset + c4_nbytes + c4_indexer_nbytes = n_c4 * c4_indexer_layer_nbytes + + c128_row_nbytes = DSV4_MLA_BYTES_PER_TOKEN + c128_rows_per_page = token_page_size // 128 + c128_layer_nbytes = c128_rows_per_page * c128_row_nbytes + c128_offset = c4_indexer_offset + c4_indexer_nbytes + c128_nbytes = n_c128 * c128_layer_nbytes + + # 128 * (576B data + 8B scale) = 74,752B; align to 576B -> 74,880B. + swa_gpu_page_nbytes = _aligned_gpu_page_nbytes( + DSV4_SWA_PAGE_SIZE, + DSV4_MLA_DATA_BYTES_PER_TOKEN, + DSV4_MLA_SCALE_BYTES, + DSV4_MLA_PAGE_ALIGN_BYTES, + ) + swa_gpu_pages_per_cpu_page = DSV4_PROMPT_CACHE_PAGE_SIZE // DSV4_SWA_PAGE_SIZE + swa_layer_nbytes = swa_gpu_pages_per_cpu_page * swa_gpu_page_nbytes + swa_offset = c128_offset + c128_nbytes + swa_nbytes = layer_num * swa_layer_nbytes + + # The c4 overlap state has two KV/score pairs, hence a 4 * dim row. + c4_state_rows = 4 + c4_state_row_nbytes = 4 * head_dim * torch._utils._element_size(torch.float32) + c4_state_offset = swa_offset + swa_nbytes + c4_state_nbytes = n_c4 * c4_state_rows * c4_state_row_nbytes + c4_indexer_state_row_nbytes = 4 * indexer_head_dim * torch._utils._element_size(torch.float32) + c4_indexer_state_offset = c4_state_offset + c4_state_nbytes + c4_indexer_state_nbytes = n_c4 * c4_state_rows * c4_indexer_state_row_nbytes + page_nbytes = c4_indexer_state_offset + c4_indexer_state_nbytes + + return cls( + token_page_size=token_page_size, + history_block_num=history_block_num, + layer_num=layer_num, + n_c4=n_c4, + n_c128=n_c128, + head_dim=head_dim, + indexer_head_dim=indexer_head_dim, + c4_offset=c4_offset, + c4_gpu_page_nbytes=c4_gpu_page_nbytes, + c4_gpu_pages_per_cpu_page=c4_gpu_pages_per_cpu_page, + c4_layer_nbytes=c4_layer_nbytes, + c4_nbytes=c4_nbytes, + c4_indexer_offset=c4_indexer_offset, + c4_indexer_gpu_page_nbytes=c4_indexer_gpu_page_nbytes, + c4_indexer_layer_nbytes=c4_indexer_layer_nbytes, + c4_indexer_nbytes=c4_indexer_nbytes, + c128_offset=c128_offset, + c128_row_nbytes=c128_row_nbytes, + c128_rows_per_page=c128_rows_per_page, + c128_layer_nbytes=c128_layer_nbytes, + c128_nbytes=c128_nbytes, + swa_offset=swa_offset, + swa_gpu_page_nbytes=swa_gpu_page_nbytes, + swa_gpu_pages_per_cpu_page=swa_gpu_pages_per_cpu_page, + swa_layer_nbytes=swa_layer_nbytes, + swa_nbytes=swa_nbytes, + c4_state_offset=c4_state_offset, + c4_state_row_nbytes=c4_state_row_nbytes, + c4_state_rows=c4_state_rows, + c4_state_nbytes=c4_state_nbytes, + c4_indexer_state_offset=c4_indexer_state_offset, + c4_indexer_state_row_nbytes=c4_indexer_state_row_nbytes, + c4_indexer_state_nbytes=c4_indexer_state_nbytes, + page_nbytes=page_nbytes, + ) + + +@dataclass(frozen=True) +class DeepseekV4CpuCacheLoadPlan: + loaded_start: int + loaded_end: int + mem_indexes: torch.Tensor + history_full_slots: torch.Tensor + history_c4_slots: Optional[torch.Tensor] + history_c128_slots: Optional[torch.Tensor] + resume_full_slots: torch.Tensor + resume_swa_slots: torch.Tensor + resume_full_slots_long: torch.Tensor + resume_swa_pages: torch.Tensor + resume_swa_page_deltas: torch.Tensor + history_c4_full_slots_long: Optional[torch.Tensor] + history_c4_pages: Optional[torch.Tensor] + history_c4_page_deltas: Optional[torch.Tensor] + history_c128_full_slots_long: Optional[torch.Tensor] + + class PackedPagePool: - """fp8_ds_mla 风格的 page-slab 存储: 每页前段连续放 token 的 data 字节,页尾放 per-token scale 字节。 + """fp8_ds_mla 风格的 packed page 存储: 每页前段连续放 token 的 data 字节,页尾放 per-token scale 字节。 寻址是纯 token 槽位 (page = slot // page_size),page 只是 scale-tail/对齐的物理打包技巧, 不存在页粒度的分配。``write``/``read`` 是 torch 参考实现(单测 oracle);生产写入走 @@ -152,6 +342,7 @@ def __init__( max_request_num: int, mtp_step: int, indexer_head_dim: int = 128, + cpu_cache_token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, @@ -174,6 +365,12 @@ def __init__( self.c4_state_ring = DSV4_C4_STATE_RING + mtp_step self.c128_state_ring = _ceil_div(DSV4_C128_STATE_RING + mtp_step, 4) * 4 self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) + self.cpu_cache_layout = DeepseekV4CpuCacheLayout.from_compress_rates( + self.compress_rates, + token_page_size=cpu_cache_token_page_size, + head_dim=head_dim, + indexer_head_dim=indexer_head_dim, + ) # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法) self.layer_to_c4_idx = {} @@ -326,6 +523,13 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): ) self._init_state_sentinel(self.c128_state_buffer) + layout = self.cpu_cache_layout + assert self.swa_pool.bytes_per_page == layout.swa_gpu_page_nbytes + if self.n_c4: + assert self.c4_pool.bytes_per_page == layout.c4_gpu_page_nbytes + assert self.c4_indexer_pool.bytes_per_page == layout.c4_indexer_gpu_page_nbytes + assert layout.c128_row_nbytes == self.mla_bytes_per_token + logger.info( f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " @@ -366,6 +570,147 @@ def get_c128_state_buffer(self, layer_index: int) -> torch.Tensor: assert self.compress_rates[layer_index] == 128, "只有 c128(HCA) 层有 request-scoped compressor state" return self.c128_state_buffer[self.layer_to_c128_idx[layer_index]] + # ------------------------------------------------------------------ CPU cache load resources + def get_loadable_cpu_cache_end( + self, + loaded_start: int, + requested_end: int, + full_token_capacity: int, + swa_page_capacity: int, + c4_page_capacity: int, + c128_slot_capacity: int, + ) -> int: + """Return the farthest loadable checkpoint boundary, or zero. + + History is independently addressable in 256-token blocks, but resume + SWA/state exists only at the end of each CPU checkpoint page. Therefore + capacity cropping must never return an intermediate 256-token boundary. + """ + loaded_start = int(loaded_start) + requested_end = int(requested_end) + page = self.cpu_cache_layout.token_page_size + if loaded_start < 0 or loaded_start % DSV4_PROMPT_CACHE_PAGE_SIZE != 0: + raise ValueError(f"DeepSeek-V4 CPU cache loaded_start must be 256-token aligned, got {loaded_start}") + if requested_end <= loaded_start or requested_end % page != 0: + raise ValueError( + f"DeepSeek-V4 CPU cache requested_end must be a checkpoint boundary after loaded_start, " + f"got start={loaded_start}, end={requested_end}, page={page}" + ) + if int(swa_page_capacity) < 2: + return 0 + + token_capacity = int(full_token_capacity) + if self.n_c4: + token_capacity = min(token_capacity, int(c4_page_capacity) * DSV4_PROMPT_CACHE_PAGE_SIZE) + if self.n_c128: + token_capacity = min(token_capacity, int(c128_slot_capacity) * 128) + token_capacity = token_capacity // DSV4_PROMPT_CACHE_PAGE_SIZE * DSV4_PROMPT_CACHE_PAGE_SIZE + loadable_end = min(requested_end, (loaded_start + token_capacity) // page * page) + return loadable_end if loadable_end > loaded_start else 0 + + def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4CpuCacheLoadPlan: + """Allocate a missing suffix and publish all derived mappings as one plan. + + ``loaded_end`` is the CPU checkpoint boundary. ``token_num`` may be + smaller than the checkpoint page when a GPU radix prefix overlaps its + beginning, but both endpoints remain 256-token aligned. + """ + token_num = int(token_num) + loaded_end = int(loaded_end) + layout = self.cpu_cache_layout + if token_num <= 0 or token_num % DSV4_PROMPT_CACHE_PAGE_SIZE != 0: + raise ValueError( + f"DeepSeek-V4 CPU cache load size must be a positive multiple of " + f"{DSV4_PROMPT_CACHE_PAGE_SIZE}, got {token_num}" + ) + if loaded_end < token_num or loaded_end % layout.token_page_size != 0: + raise ValueError( + f"DeepSeek-V4 CPU cache loaded_end must be a checkpoint boundary >= token_num, " + f"got loaded_end={loaded_end}, token_num={token_num}, page={layout.token_page_size}" + ) + + block_num = token_num // DSV4_PROMPT_CACHE_PAGE_SIZE + device = self.full_to_swa_indexs.device + full_indexes_cpu = self.alloc(token_num) + swa_pages_cpu = self._alloc_swa_pages(2) + c4_pages_cpu = self.alloc_c4_pages(block_num) if self.n_c4 else None + c128_slots_cpu = self.alloc_c128(block_num * 2) if self.n_c128 else None + + mem_indexes = full_indexes_cpu.to(device, non_blocking=True) + history_full_slots = mem_indexes.view(block_num, DSV4_PROMPT_CACHE_PAGE_SIZE) + resume_full_slots = history_full_slots[-1] + + swa_pages = swa_pages_cpu.to(device, non_blocking=True) + resume_swa_slots = ( + swa_pages[:, None] * DSV4_SWA_PAGE_SIZE + + torch.arange(DSV4_SWA_PAGE_SIZE, dtype=torch.int32, device=device)[None, :] + ).reshape(-1) + + history_c4_slots = None + if c4_pages_cpu is not None: + c4_pages = c4_pages_cpu.to(device, non_blocking=True) + history_c4_slots = ( + c4_pages[:, None] * DSV4_C4_PAGE_SIZE + + torch.arange(DSV4_C4_PAGE_SIZE, dtype=torch.int32, device=device)[None, :] + ) + + history_c128_slots = None + if c128_slots_cpu is not None: + history_c128_slots = c128_slots_cpu.to(device, non_blocking=True).view(block_num, 2) + + resume_full_slots_long = resume_full_slots.long() + resume_swa_pages = swa_pages + resume_swa_page_deltas = torch.full( + resume_swa_pages.shape, + DSV4_SWA_PAGE_SIZE, + dtype=torch.int32, + device=device, + ) + + history_c4_full_slots_long = history_c4_pages = history_c4_page_deltas = None + if history_c4_slots is not None: + history_c4_full_slots_long = history_full_slots[:, 3::4].long() + history_c4_pages = c4_pages + history_c4_page_deltas = torch.full( + history_c4_pages.shape, + DSV4_C4_PAGE_SIZE, + dtype=torch.int32, + device=device, + ) + + history_c128_full_slots_long = None + if history_c128_slots is not None: + history_c128_full_slots_long = history_full_slots[:, 127::128].long() + + return DeepseekV4CpuCacheLoadPlan( + loaded_start=loaded_end - token_num, + loaded_end=loaded_end, + mem_indexes=mem_indexes, + history_full_slots=history_full_slots, + history_c4_slots=history_c4_slots, + history_c128_slots=history_c128_slots, + resume_full_slots=resume_full_slots, + resume_swa_slots=resume_swa_slots, + resume_full_slots_long=resume_full_slots_long, + resume_swa_pages=resume_swa_pages, + resume_swa_page_deltas=resume_swa_page_deltas, + history_c4_full_slots_long=history_c4_full_slots_long, + history_c4_pages=history_c4_pages, + history_c4_page_deltas=history_c4_page_deltas, + history_c128_full_slots_long=history_c128_full_slots_long, + ) + + def commit_cpu_cache_load_plan(self, plan: DeepseekV4CpuCacheLoadPlan) -> None: + """Publish an unpacked plan without allocating at the transaction boundary.""" + self.full_to_swa_indexs[plan.resume_full_slots_long] = plan.resume_swa_slots + self.swa_page_live_count.index_add_(0, plan.resume_swa_pages, plan.resume_swa_page_deltas) + if plan.history_c4_slots is not None: + self.full_to_c4_indexs[plan.history_c4_full_slots_long] = plan.history_c4_slots + self.c4_page_live_count.index_add_(0, plan.history_c4_pages, plan.history_c4_page_deltas) + if plan.history_c128_slots is not None: + self.full_to_c128_indexs[plan.history_c128_full_slots_long] = plan.history_c128_slots + return + # ------------------------------------------------------------------ swa slot lifecycle def register_swa_free_hook(self, fn) -> None: """fn(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点 free swa)。""" @@ -741,12 +1086,12 @@ def pack_indexer_k_to_cache( ) # ------------------------------------------------------------------ fenced inherited APIs - # kv_buffer 是 page 索引的 uint8 slab,基类按 token 索引读写的接口会静默写坏数据,显式拦截。 + # kv_buffer 是 page 索引的 uint8 buffer,基类按 token 索引读写的接口会静默写坏数据,显式拦截。 def get_index_kv_buffer(self, index): - raise NotImplementedError("DeepSeek-V4 packed page-slab cache does not support token-indexed kv_buffer io") + raise NotImplementedError("DeepSeek-V4 packed page cache does not support token-indexed kv_buffer io") def load_index_kv_buffer(self, index, load_tensor_dict): - raise NotImplementedError("DeepSeek-V4 packed page-slab cache does not support token-indexed kv_buffer io") + raise NotImplementedError("DeepSeek-V4 packed page cache does not support token-indexed kv_buffer io") def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") diff --git a/lightllm/common/kv_cache_mem_manager/mem_utils.py b/lightllm/common/kv_cache_mem_manager/mem_utils.py index 79ea448794..03f1445791 100644 --- a/lightllm/common/kv_cache_mem_manager/mem_utils.py +++ b/lightllm/common/kv_cache_mem_manager/mem_utils.py @@ -21,6 +21,15 @@ def select_mem_manager_class(): # 先判断是否是 deepseek 系列的模型 model_class = get_llm_model_class() + from lightllm.models import DeepseekV4TpPartModel + + if issubclass(model_class, DeepseekV4TpPartModel): + from . import DeepseekV4MemoryManager + + mem_class = DeepseekV4MemoryManager + logger.info(f"Model kv cache using default, mem_manager class: {mem_class}") + return mem_class + from lightllm.models import Deepseek3_2TpPartModel if issubclass(model_class, Deepseek3_2TpPartModel): diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 0725ce9b93..164ea7dca0 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -93,3 +93,41 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: mem_manager: DeepseekV4MemoryManager = self.mem_manager mem_manager.pack_mla_kv_to_cache(layer_index, mem_index, kv) return + + def pack_cpu_cache_pages(self, source_mem_indexes: torch.Tensor, staging: torch.Tensor) -> None: + """Pack complete DS4 checkpoints into caller-owned CUDA staging.""" + from lightllm.models.deepseek_v4.triton_kernel.cpu_cache_io import pack_gpu_cache_to_staging + + pack_gpu_cache_to_staging(self.mem_manager, source_mem_indexes, staging) + return + + def scatter_packed_cpu_cache_pages( + self, + staging: torch.Tensor, + page_indexes: torch.Tensor, + cpu_cache_client, + ) -> None: + """Write packed checkpoints to their pinned shared-memory pages.""" + from lightllm.models.deepseek_v4.triton_kernel.cpu_cache_io import scatter_staging_to_cpu_pages + + scatter_staging_to_cpu_pages(staging, cpu_cache_client.cpu_kv_cache_tensor, page_indexes) + return + + def load_cpu_cache_pages( + self, + plan, + page_indexes: torch.Tensor, + cpu_cache_client, + first_page_history_offset_tokens: int = 0, + ) -> None: + """Selectively restore compressed history and the final resume window.""" + from lightllm.models.deepseek_v4.triton_kernel.cpu_cache_io import unpack_cpu_cache_to_gpu + + unpack_cpu_cache_to_gpu( + self.mem_manager, + plan, + cpu_cache_client.cpu_kv_cache_tensor, + page_indexes, + first_page_history_offset_tokens, + ) + return diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 64b5e361ab..073a7b3b23 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -566,6 +566,11 @@ def init_compress_state(self, req_idx: int): self.clear_runtime_state(req_idx) return + def finish_cpu_cache_load(self, req_idx: int, loaded_len: int) -> None: + """Keep only the final restored 256-token SWA page eligible for radix reuse.""" + self._swa_evict_marks[req_idx] = loaded_len - self.get_prompt_cache_page_size() + return + # ------------------------------------------------------------------ compress slot prep (per step) def _register_c4_slots(self, full_slots: torch.Tensor, slots: torch.Tensor) -> None: """写入 full->c4 槽映射并按页累加存活计数。""" diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 90ba83a70f..1022e0c03d 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -9,6 +9,7 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_CPU_CACHE_TOKEN_PAGE_SIZE from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import ( DeepseekV4PreAndPostLayerWeight, ) @@ -90,6 +91,11 @@ def _init_mem_manager(self): indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, mtp_step=self.args.mtp_step, + cpu_cache_token_page_size=( + DSV4_CPU_CACHE_TOKEN_PAGE_SIZE + if self.args.cpu_cache_token_page_size is None + else self.args.cpu_cache_token_page_size + ), mem_fraction=self.mem_fraction, ) self.req_manager.mem_manager = self.mem_manager diff --git a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py new file mode 100644 index 0000000000..224c08cb63 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py @@ -0,0 +1,765 @@ +import torch +import triton +import triton.language as tl + + +_BYTE_BLOCK = 1024 +_STATE_BLOCK = 256 +_HISTORY_BLOCK_SIZE = 256 + + +@triton.jit +def _pack_pool_gpu_pages_kernel( + full_slots, + full_to_pool, + pool, + pool_stride0, + pool_stride1, + staging, + staging_stride0, + token_page_size: tl.constexpr, + layer_num: tl.constexpr, + gpu_pages_per_cpu_page: tl.constexpr, + first_full_offset: tl.constexpr, + full_offset_per_gpu_page: tl.constexpr, + pool_page_size: tl.constexpr, + gpu_page_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + byte_blocks_per_gpu_page: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + byte_block = pid % byte_blocks_per_gpu_page + job = pid // byte_blocks_per_gpu_page + gpu_page = job % gpu_pages_per_cpu_page + job = job // gpu_pages_per_cpu_page + layer = job % layer_num + logical_page = job // layer_num + + logical_page_i64 = logical_page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + full_offset = logical_page_i64 * token_page_size + first_full_offset + gpu_page_i64 * full_offset_per_gpu_page + full_slot = tl.load(full_slots + full_offset).to(tl.int64) + pool_slot = tl.load(full_to_pool + full_slot).to(tl.int64) + physical_page = pool_slot // pool_page_size + + offsets = byte_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < gpu_page_nbytes + pool_ptr = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 + staging_ptr = ( + staging + + logical_page_i64 * staging_stride0 + + section_offset + + layer_i64 * section_layer_nbytes + + gpu_page_i64 * gpu_page_nbytes + + offsets_i64 + ) + tl.store(staging_ptr, tl.load(pool_ptr, mask=mask), mask=mask) + + +@triton.jit +def _pack_c128_rows_kernel( + full_slots, + full_to_c128, + pool, + pool_stride0, + pool_stride1, + staging, + staging_stride0, + token_page_size: tl.constexpr, + layer_num: tl.constexpr, + rows_per_page: tl.constexpr, + pool_page_size: tl.constexpr, + data_nbytes: tl.constexpr, + scale_nbytes: tl.constexpr, + scale_offset: tl.constexpr, + row_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + BLOCK: tl.constexpr, +): + job = tl.program_id(0) + row = job % rows_per_page + job = job // rows_per_page + layer = job % layer_num + logical_page = job // layer_num + + logical_page_i64 = logical_page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + row_i64 = row.to(tl.int64) + full_offset = logical_page_i64 * token_page_size + (row_i64 + 1) * 128 - 1 + full_slot = tl.load(full_slots + full_offset).to(tl.int64) + pool_slot = tl.load(full_to_c128 + full_slot).to(tl.int64) + physical_page = pool_slot // pool_page_size + token_in_page = pool_slot % pool_page_size + + offsets = tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + staging_ptr = ( + staging + + logical_page_i64 * staging_stride0 + + section_offset + + layer_i64 * section_layer_nbytes + + row_i64 * row_nbytes + + offsets_i64 + ) + pool_page = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + data_mask = offsets < data_nbytes + data_ptr = pool_page + token_in_page * data_nbytes + offsets_i64 + tl.store(staging_ptr, tl.load(data_ptr, mask=data_mask), mask=data_mask) + + scale_local = offsets_i64 - data_nbytes + scale_mask = (offsets >= data_nbytes) & (offsets < row_nbytes) + scale_ptr = pool_page + scale_offset + token_in_page * scale_nbytes + scale_local + tl.store(staging_ptr, tl.load(scale_ptr, mask=scale_mask), mask=scale_mask) + + +@triton.jit +def _pack_c4_tail_states_kernel( + full_slots, + full_to_swa, + state, + state_stride0, + state_stride1, + staging_f32, + staging_page_stride, + token_page_size: tl.constexpr, + layer_num: tl.constexpr, + swa_page_size: tl.constexpr, + state_ring: tl.constexpr, + state_width: tl.constexpr, + section_offset_f32: tl.constexpr, + section_layer_elems: tl.constexpr, + blocks_per_row: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + state_block = pid % blocks_per_row + job = pid // blocks_per_row + tail_row = job % 4 + job = job // 4 + layer = job % layer_num + logical_page = job // layer_num + + logical_page_i64 = logical_page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + tail_row_i64 = tail_row.to(tl.int64) + tail_full_offset = logical_page_i64 * token_page_size + token_page_size - 4 + tail_row_i64 + tail_full_slot = tl.load(full_slots + tail_full_offset).to(tl.int64) + tail_swa_slot = tl.load(full_to_swa + tail_full_slot).to(tl.int64) + state_row = (tail_swa_slot // swa_page_size) * state_ring + tail_swa_slot % state_ring + + offsets = state_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < state_width + state_ptr = state + layer_i64 * state_stride0 + state_row * state_stride1 + offsets_i64 + staging_ptr = ( + staging_f32 + + logical_page_i64 * staging_page_stride + + section_offset_f32 + + layer_i64 * section_layer_elems + + tail_row_i64 * state_width + + offsets_i64 + ) + tl.store(staging_ptr, tl.load(state_ptr, mask=mask), mask=mask) + + +@triton.jit +def _scatter_staging_to_cpu_kernel( + staging, + staging_stride0, + cpu_pages, + cpu_stride0, + page_indexes, + page_nbytes: tl.constexpr, + blocks_per_page: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + byte_block = pid % blocks_per_page + logical_page = pid // blocks_per_page + logical_page_i64 = logical_page.to(tl.int64) + cpu_page = tl.load(page_indexes + logical_page_i64).to(tl.int64) + + offsets = byte_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < page_nbytes + source = staging + logical_page_i64 * staging_stride0 + offsets_i64 + target = cpu_pages + cpu_page * cpu_stride0 + offsets_i64 + tl.store(target, tl.load(source, mask=mask), mask=mask, cache_modifier=".wt") + + +@triton.jit +def _unpack_history_gpu_pages_kernel( + pool_slots, + pool, + pool_stride0, + pool_stride1, + cpu_pages, + cpu_stride0, + cpu_page_indexes, + first_history_block, + layer_num: tl.constexpr, + blocks_per_cpu_page: tl.constexpr, + pool_slots_per_block: tl.constexpr, + pool_page_size: tl.constexpr, + gpu_page_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + byte_blocks_per_gpu_page: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + byte_block = pid % byte_blocks_per_gpu_page + job = pid // byte_blocks_per_gpu_page + layer = job % layer_num + history_block = job // layer_num + + history_block_i64 = history_block.to(tl.int64) + layer_i64 = layer.to(tl.int64) + absolute_block = history_block_i64 + first_history_block + cpu_page_list_index = absolute_block // blocks_per_cpu_page + block_in_cpu_page = absolute_block % blocks_per_cpu_page + cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) + + pool_slot = tl.load(pool_slots + history_block_i64 * pool_slots_per_block).to(tl.int64) + physical_page = pool_slot // pool_page_size + + offsets = byte_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < gpu_page_nbytes + source = ( + cpu_pages + + cpu_page * cpu_stride0 + + section_offset + + layer_i64 * section_layer_nbytes + + block_in_cpu_page * gpu_page_nbytes + + offsets_i64 + ) + target = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 + tl.store(target, tl.load(source, mask=mask), mask=mask) + + +@triton.jit +def _unpack_c128_rows_kernel( + c128_slots, + pool, + pool_stride0, + pool_stride1, + cpu_pages, + cpu_stride0, + cpu_page_indexes, + first_history_block, + layer_num: tl.constexpr, + blocks_per_cpu_page: tl.constexpr, + pool_page_size: tl.constexpr, + data_nbytes: tl.constexpr, + scale_nbytes: tl.constexpr, + scale_offset: tl.constexpr, + row_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + BLOCK: tl.constexpr, +): + job = tl.program_id(0) + row = job % 2 + job = job // 2 + layer = job % layer_num + history_block = job // layer_num + + history_block_i64 = history_block.to(tl.int64) + layer_i64 = layer.to(tl.int64) + row_i64 = row.to(tl.int64) + absolute_block = history_block_i64 + first_history_block + cpu_page_list_index = absolute_block // blocks_per_cpu_page + block_in_cpu_page = absolute_block % blocks_per_cpu_page + cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) + + pool_slot = tl.load(c128_slots + history_block_i64 * 2 + row_i64).to(tl.int64) + physical_page = pool_slot // pool_page_size + token_in_page = pool_slot % pool_page_size + + offsets = tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + cpu_row = block_in_cpu_page * 2 + row_i64 + source = ( + cpu_pages + + cpu_page * cpu_stride0 + + section_offset + + layer_i64 * section_layer_nbytes + + cpu_row * row_nbytes + + offsets_i64 + ) + pool_page = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + data_mask = offsets < data_nbytes + data_target = pool_page + token_in_page * data_nbytes + offsets_i64 + tl.store(data_target, tl.load(source, mask=data_mask), mask=data_mask) + + scale_local = offsets_i64 - data_nbytes + scale_mask = (offsets >= data_nbytes) & (offsets < row_nbytes) + scale_target = pool_page + scale_offset + token_in_page * scale_nbytes + scale_local + tl.store(scale_target, tl.load(source, mask=scale_mask), mask=scale_mask) + + +@triton.jit +def _unpack_resume_gpu_pages_kernel( + resume_swa_slots, + pool, + pool_stride0, + pool_stride1, + cpu_pages, + cpu_stride0, + cpu_page_indexes, + cpu_page_num, + layer_num: tl.constexpr, + gpu_pages_per_cpu_page: tl.constexpr, + pool_page_size: tl.constexpr, + gpu_page_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + byte_blocks_per_gpu_page: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + byte_block = pid % byte_blocks_per_gpu_page + job = pid // byte_blocks_per_gpu_page + gpu_page = job % gpu_pages_per_cpu_page + layer = (job // gpu_pages_per_cpu_page) % layer_num + + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + pool_slot = tl.load(resume_swa_slots + gpu_page_i64 * pool_page_size).to(tl.int64) + physical_page = pool_slot // pool_page_size + cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) + + offsets = byte_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < gpu_page_nbytes + source = ( + cpu_pages + + cpu_page * cpu_stride0 + + section_offset + + layer_i64 * section_layer_nbytes + + gpu_page_i64 * gpu_page_nbytes + + offsets_i64 + ) + target = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 + tl.store(target, tl.load(source, mask=mask), mask=mask) + + +@triton.jit +def _unpack_c4_tail_states_kernel( + resume_swa_slots, + state, + state_stride0, + state_stride1, + cpu_pages_f32, + cpu_page_stride, + cpu_page_indexes, + cpu_page_num, + layer_num: tl.constexpr, + swa_page_size: tl.constexpr, + state_ring: tl.constexpr, + state_width: tl.constexpr, + section_offset_f32: tl.constexpr, + section_layer_elems: tl.constexpr, + blocks_per_row: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + state_block = pid % blocks_per_row + job = pid // blocks_per_row + tail_row = job % 4 + layer = (job // 4) % layer_num + + layer_i64 = layer.to(tl.int64) + tail_row_i64 = tail_row.to(tl.int64) + tail_swa_slot = tl.load(resume_swa_slots + 2 * swa_page_size - 4 + tail_row_i64).to(tl.int64) + state_row = (tail_swa_slot // swa_page_size) * state_ring + tail_swa_slot % state_ring + cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) + + offsets = state_block * BLOCK + tl.arange(0, BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < state_width + source = ( + cpu_pages_f32 + + cpu_page * cpu_page_stride + + section_offset_f32 + + layer_i64 * section_layer_elems + + tail_row_i64 * state_width + + offsets_i64 + ) + target = state + layer_i64 * state_stride0 + state_row * state_stride1 + offsets_i64 + tl.store(target, tl.load(source, mask=mask), mask=mask) + + +def _launch_pack_pool_gpu_pages( + *, + full_slots, + mapping, + pool, + pool_page_size, + staging, + page_num, + token_page_size, + layer_num, + gpu_pages_per_cpu_page, + first_full_offset, + full_offset_per_gpu_page, + section_offset, + section_layer_nbytes, +): + gpu_page_nbytes = pool.shape[-1] + byte_blocks_per_gpu_page = triton.cdiv(gpu_page_nbytes, _BYTE_BLOCK) + _pack_pool_gpu_pages_kernel[(page_num * layer_num * gpu_pages_per_cpu_page * byte_blocks_per_gpu_page,)]( + full_slots, + mapping, + pool, + pool.stride(0), + pool.stride(1), + staging, + staging.stride(0), + token_page_size=token_page_size, + layer_num=layer_num, + gpu_pages_per_cpu_page=gpu_pages_per_cpu_page, + first_full_offset=first_full_offset, + full_offset_per_gpu_page=full_offset_per_gpu_page, + pool_page_size=pool_page_size, + gpu_page_nbytes=gpu_page_nbytes, + section_offset=section_offset, + section_layer_nbytes=section_layer_nbytes, + byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, + BLOCK=_BYTE_BLOCK, + num_warps=4, + ) + + +def _pack_c4_states(mem_manager, full_slots, state, staging, page_num, section_offset): + layout = mem_manager.cpu_cache_layout + state_width = state.shape[-1] + blocks_per_row = triton.cdiv(state_width, _STATE_BLOCK) + _pack_c4_tail_states_kernel[(page_num * mem_manager.n_c4 * 4 * blocks_per_row,)]( + full_slots, + mem_manager.full_to_swa_indexs, + state, + state.stride(0), + state.stride(1), + staging.view(torch.float32), + staging.stride(0) // 4, + token_page_size=layout.token_page_size, + layer_num=mem_manager.n_c4, + swa_page_size=mem_manager.swa_pool.page_size, + state_ring=mem_manager.c4_state_ring, + state_width=state_width, + section_offset_f32=section_offset // 4, + section_layer_elems=4 * state_width, + blocks_per_row=blocks_per_row, + BLOCK=_STATE_BLOCK, + num_warps=4, + ) + + +def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, staging: torch.Tensor) -> None: + """Pack a compact batch of complete CPU checkpoint pages into caller-owned CUDA staging.""" + layout = mem_manager.cpu_cache_layout + assert source_mem_indexes.is_cuda and staging.is_cuda + assert source_mem_indexes.ndim == 2 and source_mem_indexes.shape[1] == layout.token_page_size + assert staging.dtype == torch.uint8 and staging.is_contiguous() + page_num = source_mem_indexes.shape[0] + assert staging.shape == (page_num, layout.page_nbytes) + + full_slots = source_mem_indexes.reshape(-1) + if mem_manager.n_c4: + for pool, section_offset, section_layer_nbytes in ( + (mem_manager.c4_pool, layout.c4_offset, layout.c4_layer_nbytes), + ( + mem_manager.c4_indexer_pool, + layout.c4_indexer_offset, + layout.c4_indexer_layer_nbytes, + ), + ): + _launch_pack_pool_gpu_pages( + full_slots=full_slots, + mapping=mem_manager.full_to_c4_indexs, + pool=pool.buffer, + pool_page_size=pool.page_size, + staging=staging, + page_num=page_num, + token_page_size=layout.token_page_size, + layer_num=mem_manager.n_c4, + gpu_pages_per_cpu_page=layout.c4_gpu_pages_per_cpu_page, + first_full_offset=3, + full_offset_per_gpu_page=_HISTORY_BLOCK_SIZE, + section_offset=section_offset, + section_layer_nbytes=section_layer_nbytes, + ) + + if mem_manager.n_c128: + pool = mem_manager.c128_pool + _pack_c128_rows_kernel[(page_num * mem_manager.n_c128 * layout.c128_rows_per_page,)]( + full_slots, + mem_manager.full_to_c128_indexs, + pool.buffer, + pool.buffer.stride(0), + pool.buffer.stride(1), + staging, + staging.stride(0), + token_page_size=layout.token_page_size, + layer_num=mem_manager.n_c128, + rows_per_page=layout.c128_rows_per_page, + pool_page_size=pool.page_size, + data_nbytes=pool.data_bytes_per_token, + scale_nbytes=pool.scale_bytes_per_token, + scale_offset=pool.scale_offset_in_page, + row_nbytes=layout.c128_row_nbytes, + section_offset=layout.c128_offset, + section_layer_nbytes=layout.c128_layer_nbytes, + BLOCK=triton.next_power_of_2(layout.c128_row_nbytes), + num_warps=4, + ) + + _launch_pack_pool_gpu_pages( + full_slots=full_slots, + mapping=mem_manager.full_to_swa_indexs, + pool=mem_manager.swa_pool.buffer, + pool_page_size=mem_manager.swa_pool.page_size, + staging=staging, + page_num=page_num, + token_page_size=layout.token_page_size, + layer_num=mem_manager.layer_num, + gpu_pages_per_cpu_page=layout.swa_gpu_pages_per_cpu_page, + first_full_offset=layout.token_page_size - _HISTORY_BLOCK_SIZE, + full_offset_per_gpu_page=mem_manager.swa_pool.page_size, + section_offset=layout.swa_offset, + section_layer_nbytes=layout.swa_layer_nbytes, + ) + if mem_manager.n_c4: + _pack_c4_states( + mem_manager, + full_slots, + mem_manager.c4_state_buffer, + staging, + page_num, + layout.c4_state_offset, + ) + _pack_c4_states( + mem_manager, + full_slots, + mem_manager.c4_indexer_state_buffer, + staging, + page_num, + layout.c4_indexer_state_offset, + ) + + +def scatter_staging_to_cpu_pages( + staging: torch.Tensor, + cpu_pages: torch.Tensor, + page_indexes: torch.Tensor, +) -> None: + """Copy packed staging rows to their pinned, mapped shared-memory pages.""" + assert staging.is_cuda and page_indexes.is_cuda + assert staging.dtype == torch.uint8 and staging.is_contiguous() and cpu_pages.is_contiguous() + page_num, page_nbytes = staging.shape + assert page_indexes.numel() == page_num + assert cpu_pages.ndim == 2 and cpu_pages.shape[1] == page_nbytes + blocks_per_page = triton.cdiv(page_nbytes, _BYTE_BLOCK) + _scatter_staging_to_cpu_kernel[(page_num * blocks_per_page,)]( + staging, + staging.stride(0), + cpu_pages, + cpu_pages.stride(0), + page_indexes, + page_nbytes=page_nbytes, + blocks_per_page=blocks_per_page, + BLOCK=_BYTE_BLOCK, + num_warps=4, + ) + + +def _unpack_history_gpu_pages( + *, + pool_slots, + pool, + cpu_pages, + cpu_page_indexes, + first_history_block, + history_block_num, + blocks_per_cpu_page, + layer_num, + pool_slots_per_block, + section_offset, + section_layer_nbytes, +): + gpu_page_nbytes = pool.buffer.shape[-1] + byte_blocks_per_gpu_page = triton.cdiv(gpu_page_nbytes, _BYTE_BLOCK) + _unpack_history_gpu_pages_kernel[(history_block_num * layer_num * byte_blocks_per_gpu_page,)]( + pool_slots, + pool.buffer, + pool.buffer.stride(0), + pool.buffer.stride(1), + cpu_pages, + cpu_pages.stride(0), + cpu_page_indexes, + first_history_block, + layer_num=layer_num, + blocks_per_cpu_page=blocks_per_cpu_page, + pool_slots_per_block=pool_slots_per_block, + pool_page_size=pool.page_size, + gpu_page_nbytes=gpu_page_nbytes, + section_offset=section_offset, + section_layer_nbytes=section_layer_nbytes, + byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, + BLOCK=_BYTE_BLOCK, + num_warps=4, + ) + + +def _unpack_c4_states(mem_manager, resume_swa_slots, state, cpu_pages, page_indexes, section_offset): + state_width = state.shape[-1] + blocks_per_row = triton.cdiv(state_width, _STATE_BLOCK) + _unpack_c4_tail_states_kernel[(mem_manager.n_c4 * 4 * blocks_per_row,)]( + resume_swa_slots, + state, + state.stride(0), + state.stride(1), + cpu_pages.view(torch.float32), + cpu_pages.stride(0) // 4, + page_indexes, + page_indexes.numel(), + layer_num=mem_manager.n_c4, + swa_page_size=mem_manager.swa_pool.page_size, + state_ring=mem_manager.c4_state_ring, + state_width=state_width, + section_offset_f32=section_offset // 4, + section_layer_elems=4 * state_width, + blocks_per_row=blocks_per_row, + BLOCK=_STATE_BLOCK, + num_warps=4, + ) + + +def unpack_cpu_cache_to_gpu( + mem_manager, + load_plan, + cpu_pages: torch.Tensor, + cpu_page_indexes: torch.Tensor, + first_page_history_offset_tokens: int = 0, +) -> None: + """Restore missing 256-token history blocks and the final 256-token resume state. + + ``cpu_page_indexes`` names consecutive large CPU checkpoint pages. The first + page may overlap an existing radix prefix; ``first_page_history_offset_tokens`` + selects its first missing 256-token block without copying the overlapped bytes. + """ + layout = mem_manager.cpu_cache_layout + assert load_plan.history_full_slots.is_cuda and cpu_page_indexes.is_cuda + assert load_plan.history_full_slots.ndim == 2 + assert load_plan.history_full_slots.shape[1] == _HISTORY_BLOCK_SIZE + assert load_plan.resume_swa_slots.shape == (_HISTORY_BLOCK_SIZE,) + assert 0 <= first_page_history_offset_tokens < layout.token_page_size + assert first_page_history_offset_tokens % _HISTORY_BLOCK_SIZE == 0 + assert cpu_pages.ndim == 2 and cpu_pages.shape[1] == layout.page_nbytes and cpu_pages.is_contiguous() + assert cpu_page_indexes.numel() > 0 + + history_block_num = load_plan.history_full_slots.shape[0] + assert history_block_num > 0 + assert load_plan.loaded_end - load_plan.loaded_start == history_block_num * _HISTORY_BLOCK_SIZE + assert load_plan.loaded_start % layout.token_page_size == first_page_history_offset_tokens + blocks_per_cpu_page = layout.history_block_num + first_history_block = first_page_history_offset_tokens // _HISTORY_BLOCK_SIZE + expected_cpu_page_num = triton.cdiv(first_history_block + history_block_num, blocks_per_cpu_page) + assert cpu_page_indexes.numel() == expected_cpu_page_num + + if mem_manager.n_c4: + assert load_plan.history_c4_slots.shape == (history_block_num, mem_manager.c4_pool.page_size) + for pool, section_offset, section_layer_nbytes in ( + (mem_manager.c4_pool, layout.c4_offset, layout.c4_layer_nbytes), + ( + mem_manager.c4_indexer_pool, + layout.c4_indexer_offset, + layout.c4_indexer_layer_nbytes, + ), + ): + _unpack_history_gpu_pages( + pool_slots=load_plan.history_c4_slots, + pool=pool, + cpu_pages=cpu_pages, + cpu_page_indexes=cpu_page_indexes, + first_history_block=first_history_block, + history_block_num=history_block_num, + blocks_per_cpu_page=blocks_per_cpu_page, + layer_num=mem_manager.n_c4, + pool_slots_per_block=mem_manager.c4_pool.page_size, + section_offset=section_offset, + section_layer_nbytes=section_layer_nbytes, + ) + + if mem_manager.n_c128: + assert load_plan.history_c128_slots.shape == (history_block_num, 2) + pool = mem_manager.c128_pool + _unpack_c128_rows_kernel[(history_block_num * mem_manager.n_c128 * 2,)]( + load_plan.history_c128_slots, + pool.buffer, + pool.buffer.stride(0), + pool.buffer.stride(1), + cpu_pages, + cpu_pages.stride(0), + cpu_page_indexes, + first_history_block, + layer_num=mem_manager.n_c128, + blocks_per_cpu_page=blocks_per_cpu_page, + pool_page_size=pool.page_size, + data_nbytes=pool.data_bytes_per_token, + scale_nbytes=pool.scale_bytes_per_token, + scale_offset=pool.scale_offset_in_page, + row_nbytes=layout.c128_row_nbytes, + section_offset=layout.c128_offset, + section_layer_nbytes=layout.c128_layer_nbytes, + BLOCK=triton.next_power_of_2(layout.c128_row_nbytes), + num_warps=4, + ) + + pool = mem_manager.swa_pool + byte_blocks_per_gpu_page = triton.cdiv(pool.buffer.shape[-1], _BYTE_BLOCK) + _unpack_resume_gpu_pages_kernel[ + (mem_manager.layer_num * layout.swa_gpu_pages_per_cpu_page * byte_blocks_per_gpu_page,) + ]( + load_plan.resume_swa_slots, + pool.buffer, + pool.buffer.stride(0), + pool.buffer.stride(1), + cpu_pages, + cpu_pages.stride(0), + cpu_page_indexes, + cpu_page_indexes.numel(), + layer_num=mem_manager.layer_num, + gpu_pages_per_cpu_page=layout.swa_gpu_pages_per_cpu_page, + pool_page_size=pool.page_size, + gpu_page_nbytes=pool.buffer.shape[-1], + section_offset=layout.swa_offset, + section_layer_nbytes=layout.swa_layer_nbytes, + byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, + BLOCK=_BYTE_BLOCK, + num_warps=4, + ) + if mem_manager.n_c4: + _unpack_c4_states( + mem_manager, + load_plan.resume_swa_slots, + mem_manager.c4_state_buffer, + cpu_pages, + cpu_page_indexes, + layout.c4_state_offset, + ) + _unpack_c4_states( + mem_manager, + load_plan.resume_swa_slots, + mem_manager.c4_indexer_state_buffer, + cpu_pages, + cpu_page_indexes, + layout.c4_indexer_state_offset, + ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 70a786edf6..5939fc57c4 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -767,8 +767,8 @@ def make_argument_parser() -> argparse.ArgumentParser: parser.add_argument( "--cpu_cache_token_page_size", type=int, - default=256, - help="""The token page size of cpu cache""", + default=None, + help="""The token page size of cpu cache. Defaults to 2048 for DeepSeek-V4 and 256 otherwise.""", ) parser.add_argument("--enable_disk_cache", action="store_true", help="""enable disk cache to store kv cache.""") parser.add_argument( diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 40d593d0e8..585f42539d 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -27,6 +27,7 @@ has_vision_module, is_linear_att_mixed_model, auto_set_max_req_total_len, + get_model_type, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args @@ -81,6 +82,8 @@ def normal_or_p_d_start(args): auto_set_max_req_total_len(args) set_unique_server_name(args) + if args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4" and args.llm_kv_type in (None, "None"): + args.llm_kv_type = "fp8kv_dsa" if args.enable_mps: from lightllm.utils.device_utils import enable_mps @@ -304,6 +307,8 @@ def normal_or_p_d_start(args): if args.enable_cpu_cache and is_linear_att_mixed_model(args.model_dir): args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for linear hybrid att model") + elif args.enable_cpu_cache and args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = 2048 if get_model_type(args.model_dir) == "deepseek_v4" else 256 # help to manage data stored on Ceph if "s3://" in args.model_dir: @@ -536,6 +541,9 @@ def pd_master_start(args): if args.run_mode != "pd_master": return + if args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4": + raise ValueError("DeepSeek-V4 CPU cache does not support pd_master") + auto_set_max_req_total_len(args) # when use config_server to support multi pd_master node, we diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 8f2df6eba8..cc3fc9b523 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -180,7 +180,7 @@ class StartArgs: pd_node_id: int = field(default=-1) enable_cpu_cache: bool = field(default=False) cpu_cache_storage_size: float = field(default=2) - cpu_cache_token_page_size: int = field(default=64) + cpu_cache_token_page_size: Optional[int] = field(default=None) enable_disk_cache: bool = field(default=False) disk_cache_storage_size: float = field(default=10) disk_cache_dir: Optional[str] = field(default=None) diff --git a/lightllm/server/multi_level_kv_cache/cpu_cache_client.py b/lightllm/server/multi_level_kv_cache/cpu_cache_client.py index 33da63ab56..eb979d2690 100644 --- a/lightllm/server/multi_level_kv_cache/cpu_cache_client.py +++ b/lightllm/server/multi_level_kv_cache/cpu_cache_client.py @@ -1,4 +1,5 @@ import ctypes +from enum import Enum, auto from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name, get_disk_cache_prompt_limit_length from typing import List, Optional, Tuple from lightllm.utils.log_utils import init_logger @@ -10,6 +11,12 @@ logger = init_logger(__name__) +class CpuPageAllocState(Enum): + NEW_STORE_OWNER = auto() + LOADING_EXISTING = auto() + READY_EXISTING = auto() + + class CpuKvCacheClient(object): """ This class is responsible for handling cpu kv cache meta data. @@ -26,13 +33,7 @@ def __init__(self, only_create_meta_data: bool, init_shm_data: bool): if not only_create_meta_data: tensor_spec = CpuCacheTensorSpec( shm_key=self.args.cpu_kv_cache_shm_id, - shape=( - self.kv_cache_tensor_meta.page_num, - self.kv_cache_tensor_meta.layer_num, - self.kv_cache_tensor_meta.token_page_size, - self.kv_cache_tensor_meta.num_heads, - self.kv_cache_tensor_meta.get_merged_head_dim(), - ), + shape=self.kv_cache_tensor_meta.get_tensor_shape(), dtype=self.kv_cache_tensor_meta.data_type, size_bytes=self.kv_cache_tensor_meta.calcu_size(), ) @@ -68,7 +69,7 @@ def get_one_empty_page(self, hash_key: int, disk_offload_enable: bool) -> Option def allocate_one_page( self, page_items: List[_LinkedListItem], hash_key: int, disk_offload_enable: bool - ) -> Tuple[Optional[int], bool]: + ) -> Tuple[Optional[int], Optional[CpuPageAllocState]]: page_index = self.page_hash_dict.get(hash_key) if page_index is not None: page_item: _CpuPageStatus = page_items[page_index] @@ -76,39 +77,39 @@ def allocate_one_page( if page_item.ref_count == 1: page_item.del_self_from_list() if page_item.is_data_ready(): - return page_index, True + return page_index, CpuPageAllocState.READY_EXISTING else: - return page_index, False + return page_index, CpuPageAllocState.LOADING_EXISTING else: page_index = self.get_one_empty_page(hash_key=hash_key, disk_offload_enable=disk_offload_enable) if page_index is not None: - return page_index, False + return page_index, CpuPageAllocState.NEW_STORE_OWNER else: - return None, False + return None, None - def allocate_pages(self, hash_keys: List[int], disk_offload_enable: bool) -> Tuple[List[int], List[bool]]: - """ - allocate_pages will add _CpuPageStaus ref_count - """ + def allocate_pages( + self, hash_keys: List[int], disk_offload_enable: bool + ) -> Tuple[List[int], List[Optional[CpuPageAllocState]]]: + """Allocate pages and return the allocation state of each page.""" page_list = [] - ready_list = [] + state_list = [] page_items = self.page_items.linked_items for hash_key in hash_keys: - page_index, ready = self.allocate_one_page( + page_index, state = self.allocate_one_page( page_items=page_items, hash_key=hash_key, disk_offload_enable=disk_offload_enable ) if page_index is not None: page_list.append(page_index) - ready_list.append(ready) + state_list.append(state) else: page_list.append(-1) - ready_list.append(False) + state_list.append(None) break left_num = len(hash_keys) - len(page_list) page_list.extend([-1 for _ in range(left_num)]) - ready_list.extend([False for _ in range(left_num)]) - return page_list, ready_list + state_list.extend([None for _ in range(left_num)]) + return page_list, state_list def update_pages_status_to_ready( self, 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 16460bf351..60f4a16bdb 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -50,6 +50,7 @@ from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token from lightllm.server.pd_io_struct import PDChunckedTransTaskRet from .multi_level_kv_cache import MultiLevelKvCacheModule +from .dsv4_multi_level_kv_cache import Dsv4MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd @@ -256,7 +257,8 @@ def init_model(self, kvargs): self.init_mtp_draft_model(kvargs) if self.args.enable_cpu_cache: - self.multi_level_cache_module = MultiLevelKvCacheModule(self) + cache_module_cls = Dsv4MultiLevelKvCacheModule if self.is_deepseek_v4 else MultiLevelKvCacheModule + self.multi_level_cache_module = cache_module_cls(self) prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling @@ -575,7 +577,10 @@ def _get_classed_reqs( # 定期对 radix cache 进行 merge,防止查询插入的操作效率下降 self._timer_merge_radix_tree() - if self.args.enable_cpu_cache and len(g_infer_context.infer_req_ids) > 0: + if self.args.enable_cpu_cache and ( + (self.is_deepseek_v4 and self.is_master_in_dp) + or (not self.is_deepseek_v4 and len(g_infer_context.infer_req_ids) > 0) + ): self.multi_level_cache_module.update_cpu_cache_task_states() if req_ids is None: @@ -744,10 +749,13 @@ def _pre_handle_finished_reqs(self, finished_reqs: List[InferReq]): # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] + cpu_store_reqs = [] if self.args.enable_cpu_cache and self.is_deepseek_v4 and self.is_master_in_dp else None # 通用状态预先填充 is_master_in_dp = self.is_master_in_dp for req_obj in run_reqs: req_obj: InferReq = req_obj + if cpu_store_reqs is not None and req_obj.cur_kv_len < req_obj.shm_req.input_len: + cpu_store_reqs.append(req_obj) if is_chuncked_mode: new_kv_len = req_obj.get_chuncked_input_token_len() else: @@ -767,6 +775,12 @@ def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> req_obj.cur_output_len += 1 pack = InferReqUpdatePack(req_obj=req_obj, output_len=req_obj.cur_output_len) update_func_objs.append(pack) + + if cpu_store_reqs: + self.multi_level_cache_module.store_completed_prefill_pages( + reqs=cpu_store_reqs, + producer_stream=g_infer_context.get_overlap_stream(), + ) return update_func_objs # 一些可以复用的通用功能函数 diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py new file mode 100644 index 0000000000..2f058af620 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py @@ -0,0 +1,386 @@ +import bisect +import dataclasses +from collections import deque +from typing import Deque, Dict, List, Optional + +import torch +import torch.distributed as dist + +from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager +from lightllm.common.kv_cache_mem_manager.operator.deepseek import DeepseekV4MemOperator +from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuPageAllocState +from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context +from lightllm.utils.envs_utils import get_dsv4_cpu_cache_max_pages_per_task + +from .multi_level_kv_cache import MultiLevelKvCacheModule + + +class Dsv4MultiLevelKvCacheModule(MultiLevelKvCacheModule): + def __init__(self, backend): + super().__init__(backend) + self._dsv4_store_sessions: Dict[int, Dsv4CpuStoreSession] = {} + self._dsv4_store_tasks: Deque[Dsv4StoreTask] = deque() + self._dsv4_staging_slots = [Dsv4StagingSlot(), Dsv4StagingSlot()] + self._dsv4_max_pages_per_store_task = get_dsv4_cpu_cache_max_pages_per_task() + + def _try_release_dsv4_session(self, session: "Dsv4CpuStoreSession") -> None: + if not session.closing or not session.load_submitted or session.pending_task_num != 0: + return + if session.load_event is not None and not session.load_event.query(): + return + + if session.leased_pages: + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + # Cumulative hashes make the root page the most valuable entry. + # Releasing tail-to-root makes the tail oldest in the LRU. + self.cpu_cache_client.deref_pages(list(reversed(session.leased_pages))) + finally: + self.cpu_cache_client.lock.release() + del self._dsv4_store_sessions[session.request_id] + + def _poll_dsv4_store_tasks(self, wait_for_one: bool = False) -> None: + if not self._dsv4_store_tasks: + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + return + + completed = [] + if wait_for_one: + self._dsv4_store_tasks[0].store_event.synchronize() + while self._dsv4_store_tasks and self._dsv4_store_tasks[0].store_event.query(): + completed.append(self._dsv4_store_tasks.popleft()) + if completed: + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + for task in completed: + self.cpu_cache_client.update_pages_status_to_ready(task.owner_pages, deref=False) + finally: + self.cpu_cache_client.lock.release() + + touched_sessions = {} + for task in completed: + slot = self._dsv4_staging_slots[task.staging_slot] + slot.in_use = False + for session in task.sessions: + session.pending_task_num -= 1 + assert session.pending_task_num >= 0 + touched_sessions[session.request_id] = session + for session in touched_sessions.values(): + self._try_release_dsv4_session(session) + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + + def _submit_dsv4_store_batch( + self, + store_pages: List["Dsv4StorePage"], + producer_stream: torch.cuda.Stream, + ) -> None: + assert 0 < len(store_pages) <= self._dsv4_max_pages_per_store_task + sessions = {item.session.request_id: item.session for item in store_pages} + owner_pages = [item.cpu_page_index for item in store_pages] + operator: DeepseekV4MemOperator = self.backend.model.mem_manager.operator + cpu_stream = g_infer_context.get_cpu_kv_cache_stream() + + self._poll_dsv4_store_tasks() + slot_index = None + while slot_index is None: + for candidate, slot in enumerate(self._dsv4_staging_slots): + if not slot.in_use: + slot_index = candidate + break + if slot_index is None: + self._poll_dsv4_store_tasks(wait_for_one=True) + + slot = self._dsv4_staging_slots[slot_index] + page_num = len(store_pages) + page_nbytes = self.backend.model.mem_manager.cpu_cache_layout.page_nbytes + if slot.buffer is None or slot.buffer.shape[0] < page_num: + with torch.cuda.stream(producer_stream): + buffer = torch.empty((page_num, page_nbytes), dtype=torch.uint8, device="cuda") + source_mem_indexes = torch.empty( + (page_num, self.backend.model.mem_manager.cpu_cache_layout.token_page_size), + dtype=torch.int32, + device="cuda", + ) + page_indexes_cuda = torch.empty((page_num,), dtype=torch.int32, device="cuda") + page_indexes_cpu = torch.empty((page_num,), dtype=torch.int32, device="cpu", pin_memory=True) + slot.buffer = buffer + slot.source_mem_indexes = source_mem_indexes + slot.page_indexes_cuda = page_indexes_cuda + slot.page_indexes_cpu = page_indexes_cpu + slot.in_use = True + + slot.page_indexes_cpu[:page_num].numpy()[:] = owner_pages + with torch.cuda.stream(producer_stream): + source_mem_indexes = slot.source_mem_indexes[:page_num] + torch.stack([item.source_mem_indexes for item in store_pages], out=source_mem_indexes) + staging = slot.buffer[:page_num] + operator.pack_cpu_cache_pages(source_mem_indexes, staging) + pack_event = torch.cuda.Event() + pack_event.record() + + # Request free/pause runs on the current stream. It may recycle the + # original DS4 slabs after packing, but never before it. + torch.cuda.current_stream().wait_event(pack_event) + with torch.cuda.stream(cpu_stream): + cpu_stream.wait_event(pack_event) + page_indexes_cuda = slot.page_indexes_cuda[:page_num] + page_indexes_cuda.copy_(slot.page_indexes_cpu[:page_num], non_blocking=True) + operator.scatter_packed_cpu_cache_pages(staging, page_indexes_cuda, self.cpu_cache_client) + store_event = torch.cuda.Event() + store_event.record() + + for session in sessions.values(): + session.pending_task_num += 1 + self._dsv4_store_tasks.append( + Dsv4StoreTask( + owner_pages=owner_pages, + sessions=list(sessions.values()), + staging_slot=slot_index, + pack_event=pack_event, + store_event=store_event, + ) + ) + + def store_completed_prefill_pages( + self, + reqs: List[InferReq], + producer_stream: torch.cuda.Stream, + ) -> None: + """Incrementally snapshot newly completed DS4 checkpoints before source reuse.""" + layout = self.backend.model.mem_manager.cpu_cache_layout + token_page_size = layout.token_page_size + store_pages: List[Dsv4StorePage] = [] + closing_sessions = {} + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + for req in reqs: + session = self._dsv4_store_sessions.get(req.req_id) + if session is None or session.closing: + continue + token_hashes = req.shm_req.token_hash_list.get_all() + if session.disabled: + closing_sessions[session.request_id] = session + continue + if session.next_page_index >= len(token_hashes): + closing_sessions[session.request_id] = session + continue + page_lens = req.shm_req.token_hash_page_len_list.get_all() + target_page_index = bisect.bisect_right(page_lens, req.cur_kv_len) + if target_page_index <= session.next_page_index: + continue + + start_page_index = session.next_page_index + page_indexes, alloc_states = self.cpu_cache_client.allocate_pages( + token_hashes[start_page_index:target_page_index], + disk_offload_enable=False, + ) + for offset, (cpu_page_index, alloc_state) in enumerate(zip(page_indexes, alloc_states)): + if cpu_page_index == -1: + session.disabled = True + break + checkpoint_index = start_page_index + offset + session.leased_pages.append(cpu_page_index) + session.next_page_index += 1 + if alloc_state is CpuPageAllocState.NEW_STORE_OWNER: + token_start = checkpoint_index * token_page_size + source_mem_indexes = self.backend.model.req_manager.req_to_token_indexs[ + req.req_idx, token_start : token_start + token_page_size + ] + store_pages.append( + Dsv4StorePage( + session=session, + cpu_page_index=cpu_page_index, + source_mem_indexes=source_mem_indexes, + ) + ) + if session.disabled or session.next_page_index >= len(token_hashes): + closing_sessions[session.request_id] = session + finally: + self.cpu_cache_client.lock.release() + + for offset in range(0, len(store_pages), self._dsv4_max_pages_per_store_task): + self._submit_dsv4_store_batch( + store_pages[offset : offset + self._dsv4_max_pages_per_store_task], + producer_stream=producer_stream, + ) + for session in closing_sessions.values(): + session.closing = True + self._try_release_dsv4_session(session) + + def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): + idle_token_num = g_infer_context.get_can_alloc_token_num() + is_master_in_dp = self.backend.is_master_in_dp + for req in reqs: + page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() + page_len_list = req.shm_req.token_hash_page_len_list.get_all() + assert len(page_list) <= len(page_len_list) + + gpu_kv_len = int(req.cur_kv_len) + if is_master_in_dp: + session = Dsv4CpuStoreSession( + request_id=req.req_id, + next_page_index=len(page_list), + leased_pages=list(page_list), + ) + page_size = self.backend.model.mem_manager.cpu_cache_layout.token_page_size + # A radix-owned checkpoint without its CPU prefix creates an unreachable hash-chain hole. + if gpu_kv_len // page_size > session.next_page_index: + session.disabled = True + session.closing = True + self._dsv4_store_sessions[req.req_id] = session + + loaded_end = gpu_kv_len + if page_list: + mem_manager: DeepseekV4MemoryManager = self.backend.model.mem_manager + layout = mem_manager.cpu_cache_layout + requested_end = int(page_len_list[len(page_list) - 1]) + if requested_end > gpu_kv_len: + swa_capacity, c4_capacity, c128_capacity = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() + loadable_end = mem_manager.get_loadable_cpu_cache_end( + gpu_kv_len, + requested_end, + idle_token_num, + swa_capacity, + c4_capacity, + c128_capacity, + ) + if loadable_end != 0: + token_num = loadable_end - gpu_kv_len + full_need = token_num + swa_need = 2 + c4_need = token_num // 256 if mem_manager.n_c4 else 0 + c128_need = token_num // 128 if mem_manager.n_c128 else 0 + if self.backend.radix_cache is not None: + radix_cache = self.backend.radix_cache + radix_cache.free_radix_cache_to_get_enough_token(full_need) + radix_cache.free_radix_cache_to_get_enough_c4_pages(c4_need) + radix_cache.free_radix_cache_to_get_enough_c128_slots(c128_need) + swa_shortage = swa_need - int(mem_manager.swa_page_allocator.can_use_mem_size) + if swa_shortage > 0: + radix_cache.free_unreferenced_swa_pages(swa_shortage) + + loadable_end = mem_manager.get_loadable_cpu_cache_end( + gpu_kv_len, + loadable_end, + int(mem_manager.allocator.can_use_mem_size), + int(mem_manager.swa_page_allocator.can_use_mem_size), + int(mem_manager.c4_page_allocator.can_use_mem_size) if mem_manager.n_c4 else 0, + int(mem_manager.c128_allocator.can_use_mem_size) if mem_manager.n_c128 else 0, + ) + if loadable_end != 0: + loaded_end = loadable_end + token_num = loaded_end - gpu_kv_len + first_page_index = gpu_kv_len // layout.token_page_size + cpu_pages = page_list[first_page_index : loaded_end // layout.token_page_size] + page_indexes_cuda = torch.tensor(cpu_pages, dtype=torch.int32, device="cuda") + plan = mem_manager.prepare_cpu_cache_load(token_num=token_num, loaded_end=loaded_end) + mem_manager.operator.load_cpu_cache_pages( + plan=plan, + page_indexes=page_indexes_cuda, + cpu_cache_client=self.cpu_cache_client, + first_page_history_offset_tokens=gpu_kv_len % layout.token_page_size, + ) + mem_manager.commit_cpu_cache_load_plan(plan) + self.backend.model.req_manager.req_to_token_indexs[ + req.req_idx, gpu_kv_len:loaded_end + ] = plan.mem_indexes + self.backend.model.req_manager.finish_cpu_cache_load(req.req_idx, loaded_end) + req.cur_kv_len = loaded_end + idle_token_num -= token_num + + if is_master_in_dp: + req.shm_req.cpu_prompt_cache_len = loaded_end - gpu_kv_len + req.shm_req.shm_cur_kv_len = loaded_end + session.load_submitted = True + if loaded_end > gpu_kv_len: + session.load_event = torch.cuda.Event() + session.load_event.record() + + dist.barrier(group=self.init_sync_group) + if is_master_in_dp: + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + return + + def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> List[InferReq]: + if self.backend.is_master_in_dp: + for req in finished_reqs: + session = self._dsv4_store_sessions.get(req.req_id) + if session is not None: + session.closing = True + self._try_release_dsv4_session(session) + self._poll_dsv4_store_tasks() + # Source pages are fenced by the pack event. Request teardown does + # not wait for the independent staging-to-host transfer. + return finished_reqs + + def update_cpu_cache_task_states(self): + self._poll_dsv4_store_tasks() + return + + +@dataclasses.dataclass +class Dsv4CpuStoreSession: + """跟踪单个 DS4 请求持有的 CPU pages 及其异步 load/store 生命周期。""" + + request_id: int + # [0, next_page_index) 的 checkpoint pages 已处理;该值指向下一个待处理 page。 + next_page_index: int = 0 + # 本 session 持有引用的 CPU page 编号,session 释放时统一 deref。 + leased_pages: List[int] = dataclasses.field(default_factory=list) + # 已提交但尚未完成的 GPU -> CPU store batch 数量。 + pending_task_num: int = 0 + # 当前 hash 链无法继续存储,不再预留新的 CPU page。 + disabled: bool = False + # 不再接收新的 store page,等待 load/store 完成后释放 session。 + closing: bool = False + # 本请求的初始 CPU -> GPU load 流程已经提交。 + load_submitted: bool = False + # 初始 CPU -> GPU load 的完成事件;没有实际 load 时为 None。 + load_event: Optional[torch.cuda.Event] = None + + +@dataclasses.dataclass +class Dsv4StorePage: + """描述一个由当前请求负责写入的GPU page -> CPU page""" + + # 持有该 CPU page 引用并跟踪异步任务的请求 session。 + session: Dsv4CpuStoreSession + # 已预留、等待写入的目标 CPU page 编号。 + cpu_page_index: int + # 该 checkpoint page 对应的 GPU KV slot 编号。 + source_mem_indexes: torch.Tensor + + +@dataclasses.dataclass +class Dsv4StoreTask: + """跟踪一个已经提交的异步 GPU -> CPU store batch。""" + + # 本 batch 负责写入的 CPU pages,完成后统一发布为 READY。 + owner_pages: List[int] + # 本 batch 涉及的请求 session,完成后分别减少 pending_task_num。 + sessions: List[Dsv4CpuStoreSession] + # 本 batch 占用的 staging slot 编号。 + staging_slot: int + # GPU KV 已打包完成;此事件完成后原始 KV slot 可以被回收。 + pack_event: torch.cuda.Event + # staging 数据已写入 CPU;轮询此事件判断整个 batch 是否完成。 + store_event: torch.cuda.Event + + +@dataclasses.dataclass +class Dsv4StagingSlot: + """可复用的 GPU staging buffer 及其 CPU page 索引缓冲区。""" + + # 打包后的连续 GPU 字节缓冲区,形状为 [page_capacity, page_nbytes]。 + buffer: Optional[torch.Tensor] = None + # batch 内每个 checkpoint page 对应的 GPU KV slot 编号。 + source_mem_indexes: Optional[torch.Tensor] = None + # Python 写入的 pinned CPU page 编号,用于异步拷贝到 GPU。 + page_indexes_cpu: Optional[torch.Tensor] = None + # scatter kernel 使用的目标 CPU page 编号。 + page_indexes_cuda: Optional[torch.Tensor] = None + # True 表示该 slot 仍被一个未完成的 store task 占用。 + in_use: bool = False diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py index d0025a03c1..9a9d253a7c 100644 --- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py @@ -6,7 +6,10 @@ from functools import lru_cache from typing import Optional, List, Deque from collections import deque -from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuKvCacheClient +from lightllm.server.multi_level_kv_cache.cpu_cache_client import ( + CpuKvCacheClient, + CpuPageAllocState, +) from lightllm.utils.config_utils import is_linear_att_mixed_model from lightllm.utils.envs_utils import get_env_start_args from ..infer_batch import InferReq @@ -231,10 +234,11 @@ def _start_kv_cache_offload_task( try: self.cpu_cache_client.lock.acquire_sleep1ms() - page_list, ready_list = self.cpu_cache_client.allocate_pages( + page_list, alloc_states = self.cpu_cache_client.allocate_pages( token_hash_list[:move_block_size], disk_offload_enable=self.args.enable_disk_cache, ) + ready_list = [state is CpuPageAllocState.READY_EXISTING for state in alloc_states] finally: self.cpu_cache_client.lock.release() diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index bd9ce0e211..4289123a4f 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -216,6 +216,11 @@ def get_disk_cache_prompt_limit_length(): return int(os.getenv("LIGHTLLM_DISK_CACHE_PROMPT_LIMIT_LENGTH", 2048)) +@lru_cache(maxsize=None) +def get_dsv4_cpu_cache_max_pages_per_task() -> int: + return int(os.getenv("LIGHTLLM_DSV4_CPU_CACHE_MAX_PAGES_PER_TASK", 4)) + + @lru_cache(maxsize=None) def enable_huge_page(): """ diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index 494908cb10..d0116d52bc 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -15,13 +15,20 @@ get_added_mtp_kv_layer_num, ) from lightllm.utils.log_utils import init_logger -from lightllm.utils.config_utils import get_num_key_value_heads, get_head_dim, get_layer_num, is_linear_att_mixed_model +from lightllm.utils.config_utils import ( + get_config_json, + get_num_key_value_heads, + get_head_dim, + get_layer_num, + is_linear_att_mixed_model, +) from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( MemoryManager, PPLINT8KVMemoryManager, PPLINT4KVMemoryManager, Deepseek2MemoryManager, + DeepseekV4MemoryManager, Qwen3NextMemManager, ) @@ -114,11 +121,33 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": scale_head_dim=get_head_dim(args.model_dir) // 8, scale_data_type=get_llm_data_type(), ) + elif mem_manager_class is DeepseekV4MemoryManager: + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4CpuCacheLayout + + config = get_config_json(args.model_dir) + layer_num = get_layer_num(args.model_dir) + get_added_mtp_kv_layer_num() + layout = DeepseekV4CpuCacheLayout.from_compress_rates( + compress_rates=config["compress_ratios"][:layer_num], + token_page_size=args.cpu_cache_token_page_size, + head_dim=get_head_dim(args.model_dir), + indexer_head_dim=config["index_head_dim"], + ) + cpu_cache_meta = CpuKVCacheMeta( + page_num=0, + token_page_size=args.cpu_cache_token_page_size, + layer_num=layer_num, + num_heads=1, + head_dim=layout.page_nbytes, + data_type=torch.uint8, + scale_head_dim=0, + scale_data_type=torch.uint8, + page_shape=(layout.page_nbytes,), + ) else: logger.error(f"not support mem manager: {mem_manager_class} for cpu kv cache") raise Exception(f"not support mem manager: {mem_manager_class} for cpu kv cache") - if args.mtp_mode is not None: + if args.mtp_mode is not None and mem_manager_class is not DeepseekV4MemoryManager: # TODO 可能会存在不同mtp模式的精度问题 assert is_linear_att_mixed_model(args.model_dir) is False, "linear att mixed model does not support mtp mode" cpu_cache_meta.layer_num += get_added_mtp_kv_layer_num() @@ -143,11 +172,14 @@ class CpuKVCacheMeta: data_type: torch.dtype scale_head_dim: int scale_data_type: torch.dtype + page_shape: Optional[Tuple[int, ...]] = None def calcu_size(self): return self.page_num * self.calcu_one_page_size() def calcu_one_page_size(self): + if self.page_shape is not None: + return int(np.prod(self.page_shape)) * self.data_type.itemsize return ( self.token_page_size * self.layer_num @@ -155,6 +187,17 @@ def calcu_one_page_size(self): * (self.head_dim * self.data_type.itemsize + self.scale_head_dim * self.scale_data_type.itemsize) ) + def get_tensor_shape(self) -> Tuple[int, ...]: + if self.page_shape is not None: + return (self.page_num, *self.page_shape) + return ( + self.page_num, + self.layer_num, + self.token_page_size, + self.num_heads, + self.get_merged_head_dim(), + ) + def get_merged_head_dim(self): """ 返回将head_dim 和 scale_head_dim 看成融合成一个head_dim时候, head_dim的长度。 From f81bef3ec8e443814e4137b164af950f1cb4d94e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 4 Aug 2026 05:45:56 +0000 Subject: [PATCH 097/214] dsv4 support micro batch overlap --- .../fused_moe/fused_moe_weight.py | 27 ++- .../fused_moe/impl/deepgemm_impl.py | 20 +- .../layer_infer/transformer_layer_infer.py | 184 ++++++++++++++++++ 3 files changed, 228 insertions(+), 3 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index bab48b9895..aa44688a44 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -191,6 +191,23 @@ def low_latency_dispatch( scoring_func=self.scoring_func, ) + def low_latency_dispatch_with_topk( + self, + hidden_states: torch.Tensor, + topk_idx: torch.Tensor, + topk_weights: torch.Tensor, + ): + assert self.enable_ep_moe, "low_latency_dispatch_with_topk is only supported when enable_ep_moe is True" + return self.fuse_moe_impl.low_latency_dispatch_with_topk( + hidden_states=hidden_states, + topk_idx=topk_idx, + topk_weights=topk_weights, + ) + + def quantize_dispatch_input(self, hidden_states: torch.Tensor): + assert self.enable_ep_moe, "quantize_dispatch_input is only supported when enable_ep_moe is True" + return self.fuse_moe_impl.quantize_dispatch_input(hidden_states=hidden_states, w13=self.w13) + def select_experts_and_quant_input( self, hidden_states: torch.Tensor, @@ -226,7 +243,12 @@ def dispatch( ) def masked_group_gemm( - self, recv_x: Tuple[torch.Tensor], masked_m: torch.Tensor, dtype: torch.dtype, expected_m: int + self, + recv_x: Tuple[torch.Tensor], + masked_m: torch.Tensor, + dtype: torch.dtype, + expected_m: int, + clamp_limit: Optional[float] = None, ): assert self.enable_ep_moe, "masked_group_gemm is only supported when enable_ep_moe is True" return self.fuse_moe_impl.masked_group_gemm( @@ -236,6 +258,7 @@ def masked_group_gemm( masked_m=masked_m, dtype=dtype, expected_m=expected_m, + clamp_limit=clamp_limit, ) def prefilled_group_gemm( @@ -245,6 +268,7 @@ def prefilled_group_gemm( recv_topk_idx: torch.Tensor, recv_topk_weights: torch.Tensor, hidden_dtype=torch.bfloat16, + clamp_limit: Optional[float] = None, ): assert self.enable_ep_moe, "prefilled_group_gemm is only supported when enable_ep_moe is True" return self.fuse_moe_impl.prefilled_group_gemm( @@ -255,6 +279,7 @@ def prefilled_group_gemm( w13=self.w13, w2=self.w2, hidden_dtype=hidden_dtype, + clamp_limit=clamp_limit, ) def low_latency_combine( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 612c517155..aa6c0fec6e 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -118,6 +118,18 @@ def low_latency_dispatch( scoring_func=scoring_func, ) + return self.low_latency_dispatch_with_topk( + hidden_states=hidden_states, + topk_idx=topk_idx, + topk_weights=topk_weights, + ) + + def low_latency_dispatch_with_topk( + self, + hidden_states: torch.Tensor, + topk_idx: torch.Tensor, + topk_weights: torch.Tensor, + ): topk_idx = topk_idx.to(torch.long) num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() use_fp8_w8a8 = self.quant_method.method_name != "none" @@ -132,6 +144,9 @@ def low_latency_dispatch( ) return recv_x, masked_m, topk_idx, topk_weights, handle, hook + def quantize_dispatch_input(self, hidden_states: torch.Tensor, w13: WeightPack): + return quantize_fused_experts_input(hidden_states, w13, self.quant_method) + def select_experts_and_quant_input( self, hidden_states: torch.Tensor, @@ -221,6 +236,7 @@ def prefilled_group_gemm( w13: WeightPack, w2: WeightPack, hidden_dtype=torch.bfloat16, + clamp_limit: Optional[float] = None, ): device = recv_x[0].device w13_weight, w13_scale = w13.weight, w13.weight_scale @@ -273,7 +289,7 @@ def prefilled_group_gemm( # TODO fused kernel silu_out = torch.empty((all_tokens, N // 2), device=device, dtype=hidden_dtype) - silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out) + silu_and_mul_fwd(gemm_out_a.view(-1, N), silu_out, limit=clamp_limit) qsilu_out, qsilu_out_scale = per_token_group_quant_fp8( silu_out, block_size, dtype=w13_weight.dtype, column_major_scales=True, scale_tma_aligned=True ) @@ -293,7 +309,7 @@ def prefilled_group_gemm( if Autotuner.is_autotune_warmup(): _gemm_out_a = torch.zeros((1, N), device=device, dtype=hidden_dtype) _silu_out = torch.zeros((1, N // 2), device=device, dtype=hidden_dtype) - silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out) + silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out, limit=clamp_limit) _gemm_out_a, _silu_out = None, None return gather_out diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 3cc2cc1887..b60d964102 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,10 +1,13 @@ import torch +import triton from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_sm100_mega_moe from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor from lightllm.utils.vllm_utils import vllm_ops from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post @@ -137,6 +140,187 @@ def token_forward( x = self._ffn(x, infer_state, layer_weight) return self._hc_ffn_out(x, residual, post_mix, res_mix) + def overlap_tpsp_context_forward( + self, + input_embdings, + input_embdings1, + infer_state: DeepseekV4InferStateInfo, + infer_state1: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + ): + experts = layer_weight.experts_ + if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): + input_embdings = self.context_forward(input_embdings, infer_state, layer_weight) + input_embdings1 = self.context_forward(input_embdings1, infer_state1, layer_weight) + return input_embdings, input_embdings1 + + x0, residual0, post_mix0, res_mix0 = self._hc_attn_in(input_embdings, layer_weight) + x0 = self.context_attention_forward(x0, infer_state, layer_weight) + x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) + x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) + logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) + weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) + + indices0 = indices0.to(torch.long) + qinput0 = experts.quantize_dispatch_input(x0) + from deep_ep import ElasticBuffer + + dispatch_event0 = ElasticBuffer.capture() + infer_state1.call_overlap_hook() + + x1, residual1, post_mix1, res_mix1 = self._hc_attn_in(input_embdings1, layer_weight) + x1 = self.context_attention_forward(x1, infer_state1, layer_weight) + x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) + x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) + logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) + weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) + + recv_x0, recv_indices0, recv_weights0, recv_count0, handle0, dispatch_hook0 = experts.dispatch( + qinput0, + indices0, + weights0, + overlap_event=dispatch_event0, + ) + dispatch_hook0() + + indices1 = indices1.to(torch.long) + qinput1 = experts.quantize_dispatch_input(x1) + + dispatch_event1 = ElasticBuffer.capture() + recv_x1, recv_indices1, recv_weights1, recv_count1, handle1, dispatch_hook1 = experts.dispatch( + qinput1, + indices1, + weights1, + overlap_event=dispatch_event1, + ) + + moe_out0 = experts.prefilled_group_gemm( + recv_count0, + recv_x0, + recv_indices0, + recv_weights0, + hidden_dtype=x0.dtype, + clamp_limit=self.swiglu_limit, + ) + + combine_event0 = ElasticBuffer.capture() + routed0, combine_hook0 = experts.combine(moe_out0, handle0, overlap_event=combine_event0) + + dispatch_hook1() + moe_out1 = experts.prefilled_group_gemm( + recv_count1, + recv_x1, + recv_indices1, + recv_weights1, + hidden_dtype=x1.dtype, + clamp_limit=self.swiglu_limit, + ) + + combine_event1 = ElasticBuffer.capture() + routed1, combine_hook1 = experts.combine(moe_out1, handle1, overlap_event=combine_event1) + + shared0 = self._ffn_tp(x0, infer_state, layer_weight) + combine_hook0() + routed0.add_(shared0) + output0 = self._hc_ffn_out(routed0, residual0, post_mix0, res_mix0) + + shared1 = self._ffn_tp(x1, infer_state1, layer_weight) + if self.is_last_layer: + combine_hook1() + routed1.add_(shared1) + output1 = self._hc_ffn_out(routed1, residual1, post_mix1, res_mix1) + else: + + def finish(): + combine_hook1() + routed1.add_(shared1) + + infer_state1.hook = finish + output1 = routed1, residual1, post_mix1, res_mix1 + return output0, output1 + + def overlap_tpsp_token_forward( + self, + input_embdings, + input_embdings1, + infer_state: DeepseekV4InferStateInfo, + infer_state1: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + ): + experts = layer_weight.experts_ + if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): + input_embdings = self.token_forward(input_embdings, infer_state, layer_weight) + input_embdings1 = self.token_forward(input_embdings1, infer_state1, layer_weight) + return input_embdings, input_embdings1 + + x0, residual0, post_mix0, res_mix0 = self._hc_attn_in(input_embdings, layer_weight) + x0 = self.token_attention_forward(x0, infer_state, layer_weight) + x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) + x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) + logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) + weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) + infer_state1.call_overlap_hook() + + shared0 = self._ffn_tp(x0, infer_state, layer_weight) + recv_x0, masked_m0, indices0, weights0, handle0, dispatch_hook0 = experts.low_latency_dispatch_with_topk( + x0, indices0, weights0 + ) + + x1, residual1, post_mix1, res_mix1 = self._hc_attn_in(input_embdings1, layer_weight) + x1 = self.token_attention_forward(x1, infer_state1, layer_weight) + x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) + x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) + logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) + weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) + dispatch_hook0() + + shared1 = self._ffn_tp(x1, infer_state1, layer_weight) + recv_x1, masked_m1, indices1, weights1, handle1, dispatch_hook1 = experts.low_latency_dispatch_with_topk( + x1, indices1, weights1 + ) + + expected_m = triton.cdiv( + x0.shape[0] * get_global_world_size() * self.num_experts_per_tok, + layer_weight.n_routed_experts, + ) + moe_out0 = experts.masked_group_gemm( + recv_x0, + masked_m0, + x0.dtype, + expected_m, + clamp_limit=self.swiglu_limit, + ) + + dispatch_hook1() + routed0, combine_hook0 = experts.low_latency_combine(moe_out0, indices0, weights0, handle0) + + moe_out1 = experts.masked_group_gemm( + recv_x1, + masked_m1, + x1.dtype, + expected_m, + clamp_limit=self.swiglu_limit, + ) + + combine_hook0() + routed0.add_(shared0) + output0 = self._hc_ffn_out(routed0, residual0, post_mix0, res_mix0) + + routed1, combine_hook1 = experts.low_latency_combine(moe_out1, indices1, weights1, handle1) + if self.is_last_layer: + combine_hook1() + routed1.add_(shared1) + output1 = self._hc_ffn_out(routed1, residual1, post_mix1, res_mix1) + else: + + def finish(): + combine_hook1() + routed1.add_(shared1) + + infer_state1.hook = finish + output1 = routed1, residual1, post_mix1, res_mix1 + return output0, output1 + # ------------------------------------------------------------------ shared projections / cache def _select_rope(self, infer_state: DeepseekV4InferStateInfo): if self.compress_ratio: From 534c504aa8593d7f9521645a1266d4d38b8cab4f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 5 Aug 2026 06:43:10 +0000 Subject: [PATCH 098/214] support PD Disaggregation --- .../deepseek2_mem_manager.py | 6 +- .../deepseek4_mem_manager.py | 252 +++- .../kv_cache_mem_manager/mem_manager.py | 6 +- .../qwen3next_mem_manager.py | 11 +- lightllm/common/req_manager.py | 31 + .../triton_kernel/cache_staging_io.py | 271 ++++ .../deepseek_v4/triton_kernel/cpu_cache_io.py | 1226 ++++++++--------- .../deepseek_v4/triton_kernel/pd_cache_io.py | 409 ++++++ lightllm/server/api_cli.py | 8 +- lightllm/server/api_start.py | 11 +- lightllm/server/core/objs/start_args_type.py | 4 +- lightllm/server/pd_io_struct.py | 4 +- .../pd/decode_node_impl/decode_impl.py | 25 +- .../decode_node_impl/decode_trans_process.py | 17 +- .../mode_backend/pd/nccl_kv_transporter.py | 19 +- .../mode_backend/pd/nixl_kv_transporter.py | 55 +- .../pd/prefill_node_impl/prefill_impl.py | 5 +- .../prefill_trans_process.py | 9 +- 18 files changed, 1641 insertions(+), 728 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py index 9eb02b963c..78e39cba2e 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py @@ -44,6 +44,8 @@ def write_mem_to_page_kv_move_buffer( dp_index: int, mem_managers: List["MemoryManager"], dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): @@ -59,7 +61,7 @@ def write_mem_to_page_kv_move_buffer( kv_buffer=dp_mems[0].kv_buffer, mode="write", ) - return + return cur_page.numel() * cur_page.element_size() def read_page_kv_move_buffer_to_mem( self, @@ -68,6 +70,8 @@ def read_page_kv_move_buffer_to_mem( dp_index: int, mem_managers: List["MemoryManager"], dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index e56ec095a0..6494692b3e 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -50,15 +50,7 @@ def _aligned_gpu_page_nbytes(page_size: int, data_nbytes: int, scale_nbytes: int @dataclass(frozen=True) -class DeepseekV4CpuCacheLayout: - """Opaque CPU checkpoint ABI for the composite DeepSeek-V4 cache. - - Every checkpoint contains compressed history for ``token_page_size`` tokens - and only the final 256 tokens of SWA/continuation state. The history is - arranged in independent 256-token blocks so a radix hit can skip an initial - part of a checkpoint without copying or allocating it again. - """ - +class _DeepseekV4CacheLayout: token_page_size: int history_block_num: int layer_num: int @@ -69,7 +61,7 @@ class DeepseekV4CpuCacheLayout: c4_offset: int c4_gpu_page_nbytes: int - c4_gpu_pages_per_cpu_page: int + c4_gpu_pages_per_page: int c4_layer_nbytes: int c4_nbytes: int @@ -86,7 +78,7 @@ class DeepseekV4CpuCacheLayout: swa_offset: int swa_gpu_page_nbytes: int - swa_gpu_pages_per_cpu_page: int + swa_gpu_pages_per_page: int swa_layer_nbytes: int swa_nbytes: int @@ -98,27 +90,26 @@ class DeepseekV4CpuCacheLayout: c4_indexer_state_offset: int c4_indexer_state_row_nbytes: int c4_indexer_state_nbytes: int - page_nbytes: int @classmethod - def from_compress_rates( + def _history_layout( cls, compress_rates: Sequence[int], - token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, - head_dim: int = DSV4_MLA_HEAD_DIM, - indexer_head_dim: int = DSV4_INDEXER_HEAD_DIM, - ) -> "DeepseekV4CpuCacheLayout": + token_page_size: int, + head_dim: int, + indexer_head_dim: int, + ): rates = tuple(int(rate) for rate in compress_rates) if token_page_size <= 0 or token_page_size % DSV4_PROMPT_CACHE_PAGE_SIZE != 0: raise ValueError( - f"DeepSeek-V4 CPU cache token page size must be a positive multiple of " + f"DeepSeek-V4 cache page size must be a positive multiple of " f"{DSV4_PROMPT_CACHE_PAGE_SIZE}, got {token_page_size}" ) if head_dim != DSV4_MLA_HEAD_DIM: - raise ValueError(f"DeepSeek-V4 CPU cache expects head_dim={DSV4_MLA_HEAD_DIM}, got {head_dim}") + raise ValueError(f"DeepSeek-V4 cache expects head_dim={DSV4_MLA_HEAD_DIM}, got {head_dim}") if indexer_head_dim != DSV4_INDEXER_HEAD_DIM: raise ValueError( - f"DeepSeek-V4 CPU cache expects indexer_head_dim={DSV4_INDEXER_HEAD_DIM}, got {indexer_head_dim}" + f"DeepSeek-V4 cache expects indexer_head_dim={DSV4_INDEXER_HEAD_DIM}, got {indexer_head_dim}" ) layer_num = len(rates) @@ -126,25 +117,23 @@ def from_compress_rates( n_c128 = rates.count(128) history_block_num = token_page_size // DSV4_PROMPT_CACHE_PAGE_SIZE - # 64 * (576B data + 8B scale) = 37,376B; align to 576B -> 37,440B. c4_gpu_page_nbytes = _aligned_gpu_page_nbytes( DSV4_C4_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, DSV4_MLA_SCALE_BYTES, DSV4_MLA_PAGE_ALIGN_BYTES, ) - c4_gpu_pages_per_cpu_page = history_block_num - c4_layer_nbytes = c4_gpu_pages_per_cpu_page * c4_gpu_page_nbytes + c4_gpu_pages_per_page = history_block_num + c4_layer_nbytes = c4_gpu_pages_per_page * c4_gpu_page_nbytes c4_offset = 0 c4_nbytes = n_c4 * c4_layer_nbytes - # indexer_head_dim is validated as 128: 64 * (128B data + 4B scale) = 8,448B. c4_indexer_gpu_page_nbytes = _aligned_gpu_page_nbytes( DSV4_C4_PAGE_SIZE, indexer_head_dim, DSV4_INDEXER_SCALE_BYTES, ) - c4_indexer_layer_nbytes = c4_gpu_pages_per_cpu_page * c4_indexer_gpu_page_nbytes + c4_indexer_layer_nbytes = c4_gpu_pages_per_page * c4_indexer_gpu_page_nbytes c4_indexer_offset = c4_offset + c4_nbytes c4_indexer_nbytes = n_c4 * c4_indexer_layer_nbytes @@ -154,29 +143,7 @@ def from_compress_rates( c128_offset = c4_indexer_offset + c4_indexer_nbytes c128_nbytes = n_c128 * c128_layer_nbytes - # 128 * (576B data + 8B scale) = 74,752B; align to 576B -> 74,880B. - swa_gpu_page_nbytes = _aligned_gpu_page_nbytes( - DSV4_SWA_PAGE_SIZE, - DSV4_MLA_DATA_BYTES_PER_TOKEN, - DSV4_MLA_SCALE_BYTES, - DSV4_MLA_PAGE_ALIGN_BYTES, - ) - swa_gpu_pages_per_cpu_page = DSV4_PROMPT_CACHE_PAGE_SIZE // DSV4_SWA_PAGE_SIZE - swa_layer_nbytes = swa_gpu_pages_per_cpu_page * swa_gpu_page_nbytes - swa_offset = c128_offset + c128_nbytes - swa_nbytes = layer_num * swa_layer_nbytes - - # The c4 overlap state has two KV/score pairs, hence a 4 * dim row. - c4_state_rows = 4 - c4_state_row_nbytes = 4 * head_dim * torch._utils._element_size(torch.float32) - c4_state_offset = swa_offset + swa_nbytes - c4_state_nbytes = n_c4 * c4_state_rows * c4_state_row_nbytes - c4_indexer_state_row_nbytes = 4 * indexer_head_dim * torch._utils._element_size(torch.float32) - c4_indexer_state_offset = c4_state_offset + c4_state_nbytes - c4_indexer_state_nbytes = n_c4 * c4_state_rows * c4_indexer_state_row_nbytes - page_nbytes = c4_indexer_state_offset + c4_indexer_state_nbytes - - return cls( + return dict( token_page_size=token_page_size, history_block_num=history_block_num, layer_num=layer_num, @@ -186,7 +153,7 @@ def from_compress_rates( indexer_head_dim=indexer_head_dim, c4_offset=c4_offset, c4_gpu_page_nbytes=c4_gpu_page_nbytes, - c4_gpu_pages_per_cpu_page=c4_gpu_pages_per_cpu_page, + c4_gpu_pages_per_page=c4_gpu_pages_per_page, c4_layer_nbytes=c4_layer_nbytes, c4_nbytes=c4_nbytes, c4_indexer_offset=c4_indexer_offset, @@ -198,9 +165,48 @@ def from_compress_rates( c128_rows_per_page=c128_rows_per_page, c128_layer_nbytes=c128_layer_nbytes, c128_nbytes=c128_nbytes, - swa_offset=swa_offset, + swa_offset=c128_offset + c128_nbytes, + ) + + +@dataclass(frozen=True) +class DeepseekV4CpuCacheLayout(_DeepseekV4CacheLayout): + """CPU checkpoint ABI: compressed history plus the final 256-token continuation.""" + + page_nbytes: int + + @classmethod + def from_compress_rates( + cls, + compress_rates: Sequence[int], + token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, + head_dim: int = DSV4_MLA_HEAD_DIM, + indexer_head_dim: int = DSV4_INDEXER_HEAD_DIM, + ) -> "DeepseekV4CpuCacheLayout": + history = cls._history_layout(compress_rates, token_page_size, head_dim, indexer_head_dim) + swa_gpu_page_nbytes = _aligned_gpu_page_nbytes( + DSV4_SWA_PAGE_SIZE, + DSV4_MLA_DATA_BYTES_PER_TOKEN, + DSV4_MLA_SCALE_BYTES, + DSV4_MLA_PAGE_ALIGN_BYTES, + ) + swa_gpu_pages_per_page = DSV4_PROMPT_CACHE_PAGE_SIZE // DSV4_SWA_PAGE_SIZE + swa_layer_nbytes = swa_gpu_pages_per_page * swa_gpu_page_nbytes + swa_nbytes = history["layer_num"] * swa_layer_nbytes + + c4_state_rows = 4 + c4_state_row_nbytes = 4 * head_dim * torch._utils._element_size(torch.float32) + c4_state_offset = history["swa_offset"] + swa_nbytes + c4_state_nbytes = history["n_c4"] * c4_state_rows * c4_state_row_nbytes + + c4_indexer_state_row_nbytes = 4 * indexer_head_dim * torch._utils._element_size(torch.float32) + c4_indexer_state_offset = c4_state_offset + c4_state_nbytes + c4_indexer_state_nbytes = history["n_c4"] * c4_state_rows * c4_indexer_state_row_nbytes + + return cls( + **history, swa_gpu_page_nbytes=swa_gpu_page_nbytes, - swa_gpu_pages_per_cpu_page=swa_gpu_pages_per_cpu_page, + swa_gpu_pages_per_page=swa_gpu_pages_per_page, swa_layer_nbytes=swa_layer_nbytes, swa_nbytes=swa_nbytes, c4_state_offset=c4_state_offset, @@ -210,7 +216,74 @@ def from_compress_rates( c4_indexer_state_offset=c4_indexer_state_offset, c4_indexer_state_row_nbytes=c4_indexer_state_row_nbytes, c4_indexer_state_nbytes=c4_indexer_state_nbytes, - page_nbytes=page_nbytes, + page_nbytes=c4_indexer_state_offset + c4_indexer_state_nbytes, + ) + + +@dataclass(frozen=True) +class DeepseekV4PDCacheLayout(_DeepseekV4CacheLayout): + """PD page ABI: compressed history plus request-tail continuation state.""" + + c128_state_offset: int + c128_state_row_nbytes: int + c128_state_rows: int + c128_state_layer_nbytes: int + c128_state_nbytes: int + page_nbytes: int + + @classmethod + def from_compress_rates( + cls, + compress_rates: Sequence[int], + token_page_size: int, + head_dim: int = DSV4_MLA_HEAD_DIM, + indexer_head_dim: int = DSV4_INDEXER_HEAD_DIM, + ) -> "DeepseekV4PDCacheLayout": + history = cls._history_layout(compress_rates, token_page_size, head_dim, indexer_head_dim) + swa_gpu_page_nbytes = _aligned_gpu_page_nbytes( + DSV4_SWA_PAGE_SIZE, + DSV4_MLA_DATA_BYTES_PER_TOKEN, + DSV4_MLA_SCALE_BYTES, + DSV4_MLA_PAGE_ALIGN_BYTES, + ) + swa_gpu_pages_per_page = 4 + swa_layer_nbytes = swa_gpu_pages_per_page * swa_gpu_page_nbytes + swa_nbytes = history["layer_num"] * swa_layer_nbytes + + c4_state_rows = DSV4_C4_STATE_RING - 1 + c4_state_row_nbytes = 4 * head_dim * torch._utils._element_size(torch.float32) + c4_state_offset = history["swa_offset"] + swa_nbytes + c4_state_nbytes = history["n_c4"] * c4_state_rows * c4_state_row_nbytes + + c4_indexer_state_row_nbytes = 4 * indexer_head_dim * torch._utils._element_size(torch.float32) + c4_indexer_state_offset = c4_state_offset + c4_state_nbytes + c4_indexer_state_nbytes = history["n_c4"] * c4_state_rows * c4_indexer_state_row_nbytes + + c128_state_rows = DSV4_C128_STATE_RING - 1 + c128_state_row_nbytes = 2 * head_dim * torch._utils._element_size(torch.float32) + c128_state_layer_nbytes = c128_state_rows * c128_state_row_nbytes + c128_state_offset = c4_indexer_state_offset + c4_indexer_state_nbytes + c128_state_nbytes = history["n_c128"] * c128_state_layer_nbytes + + return cls( + **history, + swa_gpu_page_nbytes=swa_gpu_page_nbytes, + swa_gpu_pages_per_page=swa_gpu_pages_per_page, + swa_layer_nbytes=swa_layer_nbytes, + swa_nbytes=swa_nbytes, + c4_state_offset=c4_state_offset, + c4_state_row_nbytes=c4_state_row_nbytes, + c4_state_rows=c4_state_rows, + c4_state_nbytes=c4_state_nbytes, + c4_indexer_state_offset=c4_indexer_state_offset, + c4_indexer_state_row_nbytes=c4_indexer_state_row_nbytes, + c4_indexer_state_nbytes=c4_indexer_state_nbytes, + c128_state_offset=c128_state_offset, + c128_state_row_nbytes=c128_state_row_nbytes, + c128_state_rows=c128_state_rows, + c128_state_layer_nbytes=c128_state_layer_nbytes, + c128_state_nbytes=c128_state_nbytes, + page_nbytes=c128_state_offset + c128_state_nbytes, ) @@ -1094,10 +1167,75 @@ def load_index_kv_buffer(self, index, load_tensor_dict): raise NotImplementedError("DeepSeek-V4 packed page cache does not support token-indexed kv_buffer io") def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") + self.pd_cache_layout = DeepseekV4PDCacheLayout.from_compress_rates( + self.compress_rates, + token_page_size=page_size, + head_dim=self.head_dim, + indexer_head_dim=self.indexer_head_dim, + ) + self.kv_move_buffer = torch.empty( + (page_num, 1, 1, 1, self.pd_cache_layout.page_nbytes), + dtype=torch.uint8, + device="cuda", + ) + self._buffer_mem_indexes_tensors = [ + torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num) + ] + return self.kv_move_buffer - def write_mem_to_page_kv_move_buffer(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") + def write_mem_to_page_kv_move_buffer( + self, + mem_indexes: List[int], + page_index: int, + dp_index: int, + mem_managers: List["MemoryManager"], + dp_world_size: int, + start_kv_index: int, + request_kv_len: int, + page_kind: str = "kv", + req_idx: int = None, + ): + assert page_kind == "kv" + pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][: len(mem_indexes)] + pin_mem_indexes.numpy()[:] = mem_indexes + mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) + from lightllm.models.deepseek_v4.triton_kernel.pd_cache_io import pack_pd_cache_page + + return pack_pd_cache_page( + mem_managers[dp_index], + self.pd_cache_layout, + mem_indexes_gpu, + self.kv_move_buffer[page_index], + start_kv_index, + request_kv_len, + req_idx, + ) - def read_page_kv_move_buffer_to_mem(self, *args, **kwargs): - raise NotImplementedError("DeepSeek-V4 packed/composite KV transfer is not implemented") + def read_page_kv_move_buffer_to_mem( + self, + mem_indexes: List[int], + page_index: int, + dp_index: int, + mem_managers: List["MemoryManager"], + dp_world_size: int, + start_kv_index: int, + request_kv_len: int, + page_kind: str = "kv", + req_idx: int = None, + ): + assert page_kind == "kv" + pin_mem_indexes = self._buffer_mem_indexes_tensors[page_index][: len(mem_indexes)] + pin_mem_indexes.numpy()[:] = mem_indexes + mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) + from lightllm.models.deepseek_v4.triton_kernel.pd_cache_io import unpack_pd_cache_page + + unpack_pd_cache_page( + mem_managers[dp_index], + self.pd_cache_layout, + mem_indexes_gpu, + self.kv_move_buffer[page_index], + start_kv_index, + request_kv_len, + req_idx, + ) + return diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 4aa0937335..f473e031ad 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -113,6 +113,8 @@ def write_mem_to_page_kv_move_buffer( dp_index: int, mem_managers: List["MemoryManager"], dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): @@ -136,7 +138,7 @@ def write_mem_to_page_kv_move_buffer( # keep for debug # logger.info(f"src token tensor {self.kv_buffer[:, mem_indexes[0], 0, 0]}") # logger.info(f"src page token tensor {cur_page[0, :, 0, 0]}") - return + return cur_page.numel() * cur_page.element_size() def read_page_kv_move_buffer_to_mem( self, @@ -145,6 +147,8 @@ def read_page_kv_move_buffer_to_mem( dp_index: int, mem_managers: List["MemoryManager"], dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): 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 c7ce9d96ba..91339c6e15 100644 --- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py @@ -81,6 +81,8 @@ def write_mem_to_page_kv_move_buffer( dp_index: int, mem_managers, dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): @@ -91,6 +93,8 @@ def write_mem_to_page_kv_move_buffer( dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, + start_kv_index=start_kv_index, + request_kv_len=request_kv_len, page_kind=page_kind, req_idx=req_idx, ) @@ -99,7 +103,8 @@ def write_mem_to_page_kv_move_buffer( helper = Qwen3NextLinearAttPageHelper(self) dp_mems = helper.get_dp_mems(mem_managers, dp_index, dp_world_size) helper.write_req_to_page(page_index=page_index, req_idx=req_idx, dp_mems=dp_mems) - return + page = self.kv_move_buffer[page_index] + return page.numel() * page.element_size() def read_page_kv_move_buffer_to_mem( self, @@ -108,6 +113,8 @@ def read_page_kv_move_buffer_to_mem( dp_index: int, mem_managers, dp_world_size: int, + start_kv_index: int, + request_kv_len: int, page_kind: str = "kv", req_idx: int = None, ): @@ -118,6 +125,8 @@ def read_page_kv_move_buffer_to_mem( dp_index=dp_index, mem_managers=mem_managers, dp_world_size=dp_world_size, + start_kv_index=start_kv_index, + request_kv_len=request_kv_len, page_kind=page_kind, req_idx=req_idx, ) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 073a7b3b23..d1a907a7c5 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -514,6 +514,37 @@ def prepare_prefill( ) return + def prepare_pd_decode_cache( + self, + req_idx: int, + ready_cache_len: int, + input_len: int, + new_full_slots: torch.Tensor, + ) -> None: + """Allocate DSV4 derived slots for a suffix received by a PD decode node.""" + page = self.get_prompt_cache_page_size() + assert ready_cache_len % page == 0 + assert new_full_slots.numel() == input_len - ready_cache_len + + new_full_slots = new_full_slots.reshape(-1).to(self.req_to_token_indexs.device, non_blocking=True) + self.prepare_prefill_compress_slots( + req_list=[req_idx], + ready_list=[ready_cache_len], + seq_list=[input_len], + mem_indexes=new_full_slots, + ) + + resume_start = max(ready_cache_len, max(0, input_len // page * page - page)) + self.mem_manager.alloc_swa_prefill( + new_full_slots[resume_start - ready_cache_len :], + self.req_to_token_indexs, + req_list=[req_idx], + ready_list=[resume_start], + seq_list=[input_len], + ) + self._swa_evict_marks[req_idx] = resume_start + return + def prepare_decode_swa( self, req_list: List[int], diff --git a/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py new file mode 100644 index 0000000000..33b015947e --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py @@ -0,0 +1,271 @@ +import triton +import triton.language as tl + + +BYTE_BLOCK = 1024 +STATE_BLOCK = 256 + + +@triton.jit +def _pool_pages_kernel( + full_slots, + full_to_pool, + pool, + pool_stride0, + pool_stride1, + paired_pool, + paired_pool_stride0, + paired_pool_stride1, + full_to_c128, + c128_pool, + c128_pool_stride0, + c128_pool_stride1, + staging, + grid_layer_num: tl.constexpr, + pool_layer_num: tl.constexpr, + pool_gpu_page_num: tl.constexpr, + full_slots_page_stride: tl.constexpr, + staging_page_stride: tl.constexpr, + first_full_offset, + full_offset_per_gpu_page: tl.constexpr, + pool_page_size: tl.constexpr, + gpu_page_nbytes: tl.constexpr, + section_offset: tl.constexpr, + section_layer_nbytes: tl.constexpr, + paired_gpu_page_nbytes: tl.constexpr, + paired_section_offset: tl.constexpr, + paired_section_layer_nbytes: tl.constexpr, + c128_row_num, + c128_layer_num: tl.constexpr, + c128_first_full_offset: tl.constexpr, + c128_full_offset_per_row: tl.constexpr, + c128_pool_page_size: tl.constexpr, + c128_data_nbytes: tl.constexpr, + c128_scale_nbytes: tl.constexpr, + c128_scale_offset: tl.constexpr, + c128_row_nbytes: tl.constexpr, + c128_section_offset: tl.constexpr, + c128_section_layer_nbytes: tl.constexpr, + section_gpu_page_start, + HAS_POOL: tl.constexpr, + HAS_PAIRED_POOL: tl.constexpr, + HAS_C128: tl.constexpr, + MODE: tl.constexpr, + BLOCK: tl.constexpr, +): + job = tl.program_id(0) + gpu_page = tl.program_id(1) + byte_block = tl.program_id(2) + layer = job % grid_layer_num + logical_page = job // grid_layer_num + + logical_page_i64 = logical_page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + block_offsets = tl.arange(0, BLOCK) + offsets = byte_block * BLOCK + block_offsets + offsets_i64 = offsets.to(tl.int64) + + if HAS_POOL: + if (layer < pool_layer_num) & (gpu_page < pool_gpu_page_num): + full_slot = tl.load( + full_slots + + logical_page_i64 * full_slots_page_stride + + first_full_offset + + gpu_page_i64 * full_offset_per_gpu_page + ).to(tl.int64) + pool_slot = tl.load(full_to_pool + full_slot).to(tl.int64) + physical_page = pool_slot // pool_page_size + mask = offsets < gpu_page_nbytes + pool_ptr = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 + staging_ptr = ( + staging + + logical_page_i64 * staging_page_stride + + section_offset + + layer_i64 * section_layer_nbytes + + (section_gpu_page_start + gpu_page_i64) * gpu_page_nbytes + + offsets_i64 + ) + if MODE == 0: + tl.store(staging_ptr, tl.load(pool_ptr, mask=mask), mask=mask) + else: + tl.store(pool_ptr, tl.load(staging_ptr, mask=mask), mask=mask) + + if HAS_PAIRED_POOL: + paired_mask = offsets < paired_gpu_page_nbytes + paired_pool_ptr = ( + paired_pool + layer_i64 * paired_pool_stride0 + physical_page * paired_pool_stride1 + offsets_i64 + ) + paired_staging_ptr = ( + staging + + logical_page_i64 * staging_page_stride + + paired_section_offset + + layer_i64 * paired_section_layer_nbytes + + (section_gpu_page_start + gpu_page_i64) * paired_gpu_page_nbytes + + offsets_i64 + ) + if MODE == 0: + tl.store( + paired_staging_ptr, + tl.load(paired_pool_ptr, mask=paired_mask), + mask=paired_mask, + ) + else: + tl.store( + paired_pool_ptr, + tl.load(paired_staging_ptr, mask=paired_mask), + mask=paired_mask, + ) + + if HAS_C128: + if (byte_block < 2) & (layer < c128_layer_num): + c128_row = gpu_page * 2 + byte_block + if c128_row < c128_row_num: + c128_row_i64 = c128_row.to(tl.int64) + c128_full_slot = tl.load( + full_slots + + logical_page_i64 * full_slots_page_stride + + c128_first_full_offset + + c128_row_i64 * c128_full_offset_per_row + ).to(tl.int64) + c128_pool_slot = tl.load(full_to_c128 + c128_full_slot).to(tl.int64) + c128_physical_page = c128_pool_slot // c128_pool_page_size + c128_token_in_page = c128_pool_slot % c128_pool_page_size + c128_pool_page = c128_pool + layer_i64 * c128_pool_stride0 + c128_physical_page * c128_pool_stride1 + block_offsets_i64 = block_offsets.to(tl.int64) + c128_staging_ptr = ( + staging + + logical_page_i64 * staging_page_stride + + c128_section_offset + + layer_i64 * c128_section_layer_nbytes + + c128_row_i64 * c128_row_nbytes + + block_offsets_i64 + ) + c128_data_mask = block_offsets < c128_data_nbytes + c128_data_ptr = c128_pool_page + c128_token_in_page * c128_data_nbytes + block_offsets_i64 + if MODE == 0: + tl.store( + c128_staging_ptr, + tl.load(c128_data_ptr, mask=c128_data_mask), + mask=c128_data_mask, + ) + else: + tl.store( + c128_data_ptr, + tl.load(c128_staging_ptr, mask=c128_data_mask), + mask=c128_data_mask, + ) + + c128_scale_local = block_offsets_i64 - c128_data_nbytes + c128_scale_mask = (block_offsets >= c128_data_nbytes) & (block_offsets < c128_row_nbytes) + c128_scale_ptr = ( + c128_pool_page + c128_scale_offset + c128_token_in_page * c128_scale_nbytes + c128_scale_local + ) + if MODE == 0: + tl.store( + c128_staging_ptr, + tl.load(c128_scale_ptr, mask=c128_scale_mask), + mask=c128_scale_mask, + ) + else: + tl.store( + c128_scale_ptr, + tl.load(c128_staging_ptr, mask=c128_scale_mask), + mask=c128_scale_mask, + ) + + +def copy_pool_pages( + mode, + *, + full_slots, + mapping, + pool, + staging, + page_num, + gpu_page_num, + layer_num, + full_slots_page_stride, + staging_page_stride, + first_full_offset, + full_offset_per_gpu_page, + pool_page_size, + section_offset, + section_layer_nbytes, + section_gpu_page_start=0, + paired_pool=None, + paired_section_offset=0, + paired_section_layer_nbytes=0, + c128_mapping=None, + c128_pool=None, + c128_row_num=0, + c128_first_full_offset=0, + c128_full_offset_per_row=0, + c128_row_nbytes=0, + c128_section_offset=0, + c128_section_layer_nbytes=0, +): + """Copy the present packed pools in one launch.""" + has_pool = pool is not None + has_c128 = c128_pool is not None and c128_row_num > 0 + c128_buffer = c128_pool.buffer if has_c128 else None + c128_layer_num = c128_pool.layer_num if has_c128 else 0 + c128_pool_page_size = c128_pool.page_size if has_c128 else 0 + c128_data_nbytes = c128_pool.data_bytes_per_token if has_c128 else 0 + c128_scale_nbytes = c128_pool.scale_bytes_per_token if has_c128 else 0 + c128_scale_offset = c128_pool.scale_offset_in_page if has_c128 else 0 + grid_layer_num = max(layer_num, c128_layer_num) + grid_gpu_page_num = max(gpu_page_num, triton.cdiv(c128_row_num, 2) if has_c128 else 0) + gpu_page_nbytes = pool.shape[-1] if has_pool else 0 + paired_gpu_page_nbytes = paired_pool.shape[-1] if paired_pool is not None else 0 + byte_blocks_per_gpu_page = triton.cdiv( + max(gpu_page_nbytes, paired_gpu_page_nbytes, 2 * BYTE_BLOCK if has_c128 else 0), + BYTE_BLOCK, + ) + _pool_pages_kernel[(page_num * grid_layer_num, grid_gpu_page_num, byte_blocks_per_gpu_page)]( + full_slots, + mapping, + pool, + pool.stride(0) if has_pool else 0, + pool.stride(1) if has_pool else 0, + paired_pool, + paired_pool.stride(0) if paired_pool is not None else 0, + paired_pool.stride(1) if paired_pool is not None else 0, + c128_mapping, + c128_buffer, + c128_buffer.stride(0) if has_c128 else 0, + c128_buffer.stride(1) if has_c128 else 0, + staging, + grid_layer_num=grid_layer_num, + pool_layer_num=layer_num, + pool_gpu_page_num=gpu_page_num, + full_slots_page_stride=full_slots_page_stride, + staging_page_stride=staging_page_stride, + first_full_offset=first_full_offset, + full_offset_per_gpu_page=full_offset_per_gpu_page, + pool_page_size=pool_page_size, + gpu_page_nbytes=gpu_page_nbytes, + section_offset=section_offset, + section_layer_nbytes=section_layer_nbytes, + paired_gpu_page_nbytes=paired_gpu_page_nbytes, + paired_section_offset=paired_section_offset, + paired_section_layer_nbytes=paired_section_layer_nbytes, + c128_row_num=c128_row_num, + c128_layer_num=c128_layer_num, + c128_first_full_offset=c128_first_full_offset, + c128_full_offset_per_row=c128_full_offset_per_row, + c128_pool_page_size=c128_pool_page_size, + c128_data_nbytes=c128_data_nbytes, + c128_scale_nbytes=c128_scale_nbytes, + c128_scale_offset=c128_scale_offset, + c128_row_nbytes=c128_row_nbytes, + c128_section_offset=c128_section_offset, + c128_section_layer_nbytes=c128_section_layer_nbytes, + section_gpu_page_start=section_gpu_page_start, + HAS_POOL=has_pool, + HAS_PAIRED_POOL=paired_pool is not None, + HAS_C128=has_c128, + MODE=0 if mode == "pack" else 1, + BLOCK=BYTE_BLOCK, + num_warps=4, + ) diff --git a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py index 224c08cb63..724c1afed7 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py @@ -2,169 +2,10 @@ import triton import triton.language as tl +from .cache_staging_io import BYTE_BLOCK as _BYTE_BLOCK -_BYTE_BLOCK = 1024 -_STATE_BLOCK = 256 _HISTORY_BLOCK_SIZE = 256 - - -@triton.jit -def _pack_pool_gpu_pages_kernel( - full_slots, - full_to_pool, - pool, - pool_stride0, - pool_stride1, - staging, - staging_stride0, - token_page_size: tl.constexpr, - layer_num: tl.constexpr, - gpu_pages_per_cpu_page: tl.constexpr, - first_full_offset: tl.constexpr, - full_offset_per_gpu_page: tl.constexpr, - pool_page_size: tl.constexpr, - gpu_page_nbytes: tl.constexpr, - section_offset: tl.constexpr, - section_layer_nbytes: tl.constexpr, - byte_blocks_per_gpu_page: tl.constexpr, - BLOCK: tl.constexpr, -): - pid = tl.program_id(0) - byte_block = pid % byte_blocks_per_gpu_page - job = pid // byte_blocks_per_gpu_page - gpu_page = job % gpu_pages_per_cpu_page - job = job // gpu_pages_per_cpu_page - layer = job % layer_num - logical_page = job // layer_num - - logical_page_i64 = logical_page.to(tl.int64) - layer_i64 = layer.to(tl.int64) - gpu_page_i64 = gpu_page.to(tl.int64) - full_offset = logical_page_i64 * token_page_size + first_full_offset + gpu_page_i64 * full_offset_per_gpu_page - full_slot = tl.load(full_slots + full_offset).to(tl.int64) - pool_slot = tl.load(full_to_pool + full_slot).to(tl.int64) - physical_page = pool_slot // pool_page_size - - offsets = byte_block * BLOCK + tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - mask = offsets < gpu_page_nbytes - pool_ptr = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 - staging_ptr = ( - staging - + logical_page_i64 * staging_stride0 - + section_offset - + layer_i64 * section_layer_nbytes - + gpu_page_i64 * gpu_page_nbytes - + offsets_i64 - ) - tl.store(staging_ptr, tl.load(pool_ptr, mask=mask), mask=mask) - - -@triton.jit -def _pack_c128_rows_kernel( - full_slots, - full_to_c128, - pool, - pool_stride0, - pool_stride1, - staging, - staging_stride0, - token_page_size: tl.constexpr, - layer_num: tl.constexpr, - rows_per_page: tl.constexpr, - pool_page_size: tl.constexpr, - data_nbytes: tl.constexpr, - scale_nbytes: tl.constexpr, - scale_offset: tl.constexpr, - row_nbytes: tl.constexpr, - section_offset: tl.constexpr, - section_layer_nbytes: tl.constexpr, - BLOCK: tl.constexpr, -): - job = tl.program_id(0) - row = job % rows_per_page - job = job // rows_per_page - layer = job % layer_num - logical_page = job // layer_num - - logical_page_i64 = logical_page.to(tl.int64) - layer_i64 = layer.to(tl.int64) - row_i64 = row.to(tl.int64) - full_offset = logical_page_i64 * token_page_size + (row_i64 + 1) * 128 - 1 - full_slot = tl.load(full_slots + full_offset).to(tl.int64) - pool_slot = tl.load(full_to_c128 + full_slot).to(tl.int64) - physical_page = pool_slot // pool_page_size - token_in_page = pool_slot % pool_page_size - - offsets = tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - staging_ptr = ( - staging - + logical_page_i64 * staging_stride0 - + section_offset - + layer_i64 * section_layer_nbytes - + row_i64 * row_nbytes - + offsets_i64 - ) - pool_page = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 - data_mask = offsets < data_nbytes - data_ptr = pool_page + token_in_page * data_nbytes + offsets_i64 - tl.store(staging_ptr, tl.load(data_ptr, mask=data_mask), mask=data_mask) - - scale_local = offsets_i64 - data_nbytes - scale_mask = (offsets >= data_nbytes) & (offsets < row_nbytes) - scale_ptr = pool_page + scale_offset + token_in_page * scale_nbytes + scale_local - tl.store(staging_ptr, tl.load(scale_ptr, mask=scale_mask), mask=scale_mask) - - -@triton.jit -def _pack_c4_tail_states_kernel( - full_slots, - full_to_swa, - state, - state_stride0, - state_stride1, - staging_f32, - staging_page_stride, - token_page_size: tl.constexpr, - layer_num: tl.constexpr, - swa_page_size: tl.constexpr, - state_ring: tl.constexpr, - state_width: tl.constexpr, - section_offset_f32: tl.constexpr, - section_layer_elems: tl.constexpr, - blocks_per_row: tl.constexpr, - BLOCK: tl.constexpr, -): - pid = tl.program_id(0) - state_block = pid % blocks_per_row - job = pid // blocks_per_row - tail_row = job % 4 - job = job // 4 - layer = job % layer_num - logical_page = job // layer_num - - logical_page_i64 = logical_page.to(tl.int64) - layer_i64 = layer.to(tl.int64) - tail_row_i64 = tail_row.to(tl.int64) - tail_full_offset = logical_page_i64 * token_page_size + token_page_size - 4 + tail_row_i64 - tail_full_slot = tl.load(full_slots + tail_full_offset).to(tl.int64) - tail_swa_slot = tl.load(full_to_swa + tail_full_slot).to(tl.int64) - state_row = (tail_swa_slot // swa_page_size) * state_ring + tail_swa_slot % state_ring - - offsets = state_block * BLOCK + tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - mask = offsets < state_width - state_ptr = state + layer_i64 * state_stride0 + state_row * state_stride1 + offsets_i64 - staging_ptr = ( - staging_f32 - + logical_page_i64 * staging_page_stride - + section_offset_f32 - + layer_i64 * section_layer_elems - + tail_row_i64 * state_width - + offsets_i64 - ) - tl.store(staging_ptr, tl.load(state_ptr, mask=mask), mask=mask) +_C128_RATIO = 128 @triton.jit @@ -193,273 +34,451 @@ def _scatter_staging_to_cpu_kernel( @triton.jit -def _unpack_history_gpu_pages_kernel( - pool_slots, - pool, - pool_stride0, - pool_stride1, - cpu_pages, - cpu_stride0, - cpu_page_indexes, - first_history_block, +def _pack_gpu_cache_to_staging_kernel( + full_slots, + full_to_c4, + c4_pool, + c4_pool_stride0, + c4_pool_stride1, + c4_indexer_pool, + c4_indexer_pool_stride0, + c4_indexer_pool_stride1, + full_to_c128, + c128_pool, + c128_pool_stride0, + c128_pool_stride1, + full_to_swa, + swa_pool, + swa_pool_stride0, + swa_pool_stride1, + c4_state, + c4_state_stride0, + c4_state_stride1, + c4_indexer_state, + c4_indexer_state_stride0, + c4_indexer_state_stride1, + staging, + staging_stride0, + programs_per_page: tl.constexpr, + c4_program_num: tl.constexpr, + c128_program_num: tl.constexpr, + swa_program_num: tl.constexpr, + token_page_size: tl.constexpr, + history_block_size: tl.constexpr, + c128_ratio: tl.constexpr, + c4_layer_num: tl.constexpr, + c4_gpu_page_num: tl.constexpr, + c4_pool_page_size: tl.constexpr, + c4_gpu_page_nbytes: tl.constexpr, + c4_section_offset: tl.constexpr, + c4_section_layer_nbytes: tl.constexpr, + c4_indexer_gpu_page_nbytes: tl.constexpr, + c4_indexer_section_offset: tl.constexpr, + c4_indexer_section_layer_nbytes: tl.constexpr, + c4_blocks_per_gpu_page: tl.constexpr, + c128_layer_num: tl.constexpr, + c128_row_num: tl.constexpr, + c128_pool_page_size: tl.constexpr, + c128_data_nbytes: tl.constexpr, + c128_scale_nbytes: tl.constexpr, + c128_scale_offset: tl.constexpr, + c128_row_nbytes: tl.constexpr, + c128_section_offset: tl.constexpr, + c128_section_layer_nbytes: tl.constexpr, layer_num: tl.constexpr, - blocks_per_cpu_page: tl.constexpr, - pool_slots_per_block: tl.constexpr, - pool_page_size: tl.constexpr, - gpu_page_nbytes: tl.constexpr, - section_offset: tl.constexpr, - section_layer_nbytes: tl.constexpr, - byte_blocks_per_gpu_page: tl.constexpr, + swa_gpu_page_num: tl.constexpr, + swa_pool_page_size: tl.constexpr, + swa_gpu_page_nbytes: tl.constexpr, + swa_section_offset: tl.constexpr, + swa_section_layer_nbytes: tl.constexpr, + swa_blocks_per_gpu_page: tl.constexpr, + c4_state_ring: tl.constexpr, + c4_state_row_nbytes: tl.constexpr, + c4_state_section_offset: tl.constexpr, + c4_indexer_state_row_nbytes: tl.constexpr, + c4_indexer_state_section_offset: tl.constexpr, + c4_state_blocks_per_row: tl.constexpr, + HAS_C4: tl.constexpr, + HAS_C128: tl.constexpr, BLOCK: tl.constexpr, ): pid = tl.program_id(0) - byte_block = pid % byte_blocks_per_gpu_page - job = pid // byte_blocks_per_gpu_page - layer = job % layer_num - history_block = job // layer_num - - history_block_i64 = history_block.to(tl.int64) - layer_i64 = layer.to(tl.int64) - absolute_block = history_block_i64 + first_history_block - cpu_page_list_index = absolute_block // blocks_per_cpu_page - block_in_cpu_page = absolute_block % blocks_per_cpu_page - cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) - - pool_slot = tl.load(pool_slots + history_block_i64 * pool_slots_per_block).to(tl.int64) - physical_page = pool_slot // pool_page_size - - offsets = byte_block * BLOCK + tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - mask = offsets < gpu_page_nbytes - source = ( - cpu_pages - + cpu_page * cpu_stride0 - + section_offset - + layer_i64 * section_layer_nbytes - + block_in_cpu_page * gpu_page_nbytes - + offsets_i64 - ) - target = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 - tl.store(target, tl.load(source, mask=mask), mask=mask) - - -@triton.jit -def _unpack_c128_rows_kernel( - c128_slots, - pool, - pool_stride0, - pool_stride1, - cpu_pages, - cpu_stride0, - cpu_page_indexes, - first_history_block, - layer_num: tl.constexpr, - blocks_per_cpu_page: tl.constexpr, - pool_page_size: tl.constexpr, - data_nbytes: tl.constexpr, - scale_nbytes: tl.constexpr, - scale_offset: tl.constexpr, - row_nbytes: tl.constexpr, - section_offset: tl.constexpr, - section_layer_nbytes: tl.constexpr, - BLOCK: tl.constexpr, -): - job = tl.program_id(0) - row = job % 2 - job = job // 2 - layer = job % layer_num - history_block = job // layer_num - - history_block_i64 = history_block.to(tl.int64) - layer_i64 = layer.to(tl.int64) - row_i64 = row.to(tl.int64) - absolute_block = history_block_i64 + first_history_block - cpu_page_list_index = absolute_block // blocks_per_cpu_page - block_in_cpu_page = absolute_block % blocks_per_cpu_page - cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) - - pool_slot = tl.load(c128_slots + history_block_i64 * 2 + row_i64).to(tl.int64) - physical_page = pool_slot // pool_page_size - token_in_page = pool_slot % pool_page_size - - offsets = tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - cpu_row = block_in_cpu_page * 2 + row_i64 - source = ( - cpu_pages - + cpu_page * cpu_stride0 - + section_offset - + layer_i64 * section_layer_nbytes - + cpu_row * row_nbytes - + offsets_i64 - ) - pool_page = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 - data_mask = offsets < data_nbytes - data_target = pool_page + token_in_page * data_nbytes + offsets_i64 - tl.store(data_target, tl.load(source, mask=data_mask), mask=data_mask) - - scale_local = offsets_i64 - data_nbytes - scale_mask = (offsets >= data_nbytes) & (offsets < row_nbytes) - scale_target = pool_page + scale_offset + token_in_page * scale_nbytes + scale_local - tl.store(scale_target, tl.load(source, mask=scale_mask), mask=scale_mask) + logical_page = pid // programs_per_page + page_pid = pid % programs_per_page + logical_page_i64 = logical_page.to(tl.int64) + offsets_base = tl.arange(0, BLOCK) + + if HAS_C4: + if page_pid < c4_program_num: + byte_block = page_pid % c4_blocks_per_gpu_page + job = page_pid // c4_blocks_per_gpu_page + gpu_page = job % c4_gpu_page_num + layer = job // c4_gpu_page_num + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + full_slot = tl.load( + full_slots + logical_page_i64 * token_page_size + 3 + gpu_page_i64 * history_block_size + ).to(tl.int64) + pool_slot = tl.load(full_to_c4 + full_slot).to(tl.int64) + physical_page = pool_slot // c4_pool_page_size + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + + c4_mask = offsets < c4_gpu_page_nbytes + c4_source = c4_pool + layer_i64 * c4_pool_stride0 + physical_page * c4_pool_stride1 + offsets_i64 + c4_target = ( + staging + + logical_page_i64 * staging_stride0 + + c4_section_offset + + layer_i64 * c4_section_layer_nbytes + + gpu_page_i64 * c4_gpu_page_nbytes + + offsets_i64 + ) + tl.store(c4_target, tl.load(c4_source, mask=c4_mask), mask=c4_mask) + + indexer_mask = offsets < c4_indexer_gpu_page_nbytes + indexer_source = ( + c4_indexer_pool + + layer_i64 * c4_indexer_pool_stride0 + + physical_page * c4_indexer_pool_stride1 + + offsets_i64 + ) + indexer_target = ( + staging + + logical_page_i64 * staging_stride0 + + c4_indexer_section_offset + + layer_i64 * c4_indexer_section_layer_nbytes + + gpu_page_i64 * c4_indexer_gpu_page_nbytes + + offsets_i64 + ) + tl.store(indexer_target, tl.load(indexer_source, mask=indexer_mask), mask=indexer_mask) + + c128_pid = page_pid - c4_program_num + if HAS_C128: + if (c128_pid >= 0) & (c128_pid < c128_program_num): + row = c128_pid % c128_row_num + layer = c128_pid // c128_row_num + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_slot = tl.load( + full_slots + logical_page_i64 * token_page_size + c128_ratio - 1 + row_i64 * c128_ratio + ).to(tl.int64) + pool_slot = tl.load(full_to_c128 + full_slot).to(tl.int64) + physical_page = pool_slot // c128_pool_page_size + token_in_page = pool_slot % c128_pool_page_size + offsets_i64 = offsets_base.to(tl.int64) + target = ( + staging + + logical_page_i64 * staging_stride0 + + c128_section_offset + + layer_i64 * c128_section_layer_nbytes + + row_i64 * c128_row_nbytes + + offsets_i64 + ) + pool_page = c128_pool + layer_i64 * c128_pool_stride0 + physical_page * c128_pool_stride1 + data_mask = offsets_base < c128_data_nbytes + data_source = pool_page + token_in_page * c128_data_nbytes + offsets_i64 + tl.store(target, tl.load(data_source, mask=data_mask), mask=data_mask) + + scale_local = offsets_i64 - c128_data_nbytes + scale_mask = (offsets_base >= c128_data_nbytes) & (offsets_base < c128_row_nbytes) + scale_source = pool_page + c128_scale_offset + token_in_page * c128_scale_nbytes + scale_local + tl.store(target, tl.load(scale_source, mask=scale_mask), mask=scale_mask) + + swa_pid = c128_pid - c128_program_num + if (swa_pid >= 0) & (swa_pid < swa_program_num): + byte_block = swa_pid % swa_blocks_per_gpu_page + job = swa_pid // swa_blocks_per_gpu_page + gpu_page = job % swa_gpu_page_num + layer = job // swa_gpu_page_num + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + full_slot = tl.load( + full_slots + + logical_page_i64 * token_page_size + + token_page_size + - history_block_size + + gpu_page_i64 * swa_pool_page_size + ).to(tl.int64) + pool_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + physical_page = pool_slot // swa_pool_page_size + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + mask = offsets < swa_gpu_page_nbytes + source = swa_pool + layer_i64 * swa_pool_stride0 + physical_page * swa_pool_stride1 + offsets_i64 + target = ( + staging + + logical_page_i64 * staging_stride0 + + swa_section_offset + + layer_i64 * swa_section_layer_nbytes + + gpu_page_i64 * swa_gpu_page_nbytes + + offsets_i64 + ) + tl.store(target, tl.load(source, mask=mask), mask=mask) + + state_pid = swa_pid - swa_program_num + if HAS_C4: + if state_pid >= 0: + byte_block = state_pid % c4_state_blocks_per_row + job = state_pid // c4_state_blocks_per_row + row = job % 4 + layer = job // 4 + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_slot = tl.load(full_slots + logical_page_i64 * token_page_size + token_page_size - 4 + row_i64).to( + tl.int64 + ) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + state_row = (swa_slot // swa_pool_page_size) * c4_state_ring + swa_slot % c4_state_ring + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + + state_mask = offsets < c4_state_row_nbytes + state_source = c4_state + layer_i64 * c4_state_stride0 + state_row * c4_state_stride1 + offsets_i64 + state_target = ( + staging + + logical_page_i64 * staging_stride0 + + c4_state_section_offset + + (layer_i64 * 4 + row_i64) * c4_state_row_nbytes + + offsets_i64 + ) + tl.store(state_target, tl.load(state_source, mask=state_mask), mask=state_mask) + + indexer_mask = offsets < c4_indexer_state_row_nbytes + indexer_source = ( + c4_indexer_state + + layer_i64 * c4_indexer_state_stride0 + + state_row * c4_indexer_state_stride1 + + offsets_i64 + ) + indexer_target = ( + staging + + logical_page_i64 * staging_stride0 + + c4_indexer_state_section_offset + + (layer_i64 * 4 + row_i64) * c4_indexer_state_row_nbytes + + offsets_i64 + ) + tl.store(indexer_target, tl.load(indexer_source, mask=indexer_mask), mask=indexer_mask) @triton.jit -def _unpack_resume_gpu_pages_kernel( +def _unpack_cpu_cache_to_gpu_kernel( + history_c4_slots, + history_c128_slots, resume_swa_slots, - pool, - pool_stride0, - pool_stride1, + c4_pool, + c4_pool_stride0, + c4_pool_stride1, + c4_indexer_pool, + c4_indexer_pool_stride0, + c4_indexer_pool_stride1, + c128_pool, + c128_pool_stride0, + c128_pool_stride1, + swa_pool, + swa_pool_stride0, + swa_pool_stride1, + c4_state, + c4_state_stride0, + c4_state_stride1, + c4_indexer_state, + c4_indexer_state_stride0, + c4_indexer_state_stride1, cpu_pages, cpu_stride0, cpu_page_indexes, cpu_page_num, + first_history_block, + c4_program_num, + c128_program_num, + c4_layer_num: tl.constexpr, + c4_pool_page_size: tl.constexpr, + c4_gpu_page_nbytes: tl.constexpr, + c4_section_offset: tl.constexpr, + c4_section_layer_nbytes: tl.constexpr, + c4_indexer_gpu_page_nbytes: tl.constexpr, + c4_indexer_section_offset: tl.constexpr, + c4_indexer_section_layer_nbytes: tl.constexpr, + c4_blocks_per_gpu_page: tl.constexpr, + c128_layer_num: tl.constexpr, + c128_pool_page_size: tl.constexpr, + c128_data_nbytes: tl.constexpr, + c128_scale_nbytes: tl.constexpr, + c128_scale_offset: tl.constexpr, + c128_row_nbytes: tl.constexpr, + c128_section_offset: tl.constexpr, + c128_section_layer_nbytes: tl.constexpr, + blocks_per_cpu_page: tl.constexpr, layer_num: tl.constexpr, - gpu_pages_per_cpu_page: tl.constexpr, - pool_page_size: tl.constexpr, - gpu_page_nbytes: tl.constexpr, - section_offset: tl.constexpr, - section_layer_nbytes: tl.constexpr, - byte_blocks_per_gpu_page: tl.constexpr, - BLOCK: tl.constexpr, -): - pid = tl.program_id(0) - byte_block = pid % byte_blocks_per_gpu_page - job = pid // byte_blocks_per_gpu_page - gpu_page = job % gpu_pages_per_cpu_page - layer = (job // gpu_pages_per_cpu_page) % layer_num - - layer_i64 = layer.to(tl.int64) - gpu_page_i64 = gpu_page.to(tl.int64) - pool_slot = tl.load(resume_swa_slots + gpu_page_i64 * pool_page_size).to(tl.int64) - physical_page = pool_slot // pool_page_size - cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) - - offsets = byte_block * BLOCK + tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - mask = offsets < gpu_page_nbytes - source = ( - cpu_pages - + cpu_page * cpu_stride0 - + section_offset - + layer_i64 * section_layer_nbytes - + gpu_page_i64 * gpu_page_nbytes - + offsets_i64 - ) - target = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 - tl.store(target, tl.load(source, mask=mask), mask=mask) - - -@triton.jit -def _unpack_c4_tail_states_kernel( - resume_swa_slots, - state, - state_stride0, - state_stride1, - cpu_pages_f32, - cpu_page_stride, - cpu_page_indexes, - cpu_page_num, - layer_num: tl.constexpr, - swa_page_size: tl.constexpr, - state_ring: tl.constexpr, - state_width: tl.constexpr, - section_offset_f32: tl.constexpr, - section_layer_elems: tl.constexpr, - blocks_per_row: tl.constexpr, + swa_program_num: tl.constexpr, + swa_gpu_page_num: tl.constexpr, + swa_pool_page_size: tl.constexpr, + swa_gpu_page_nbytes: tl.constexpr, + swa_section_offset: tl.constexpr, + swa_section_layer_nbytes: tl.constexpr, + swa_blocks_per_gpu_page: tl.constexpr, + c4_state_ring: tl.constexpr, + c4_state_row_nbytes: tl.constexpr, + c4_state_section_offset: tl.constexpr, + c4_indexer_state_row_nbytes: tl.constexpr, + c4_indexer_state_section_offset: tl.constexpr, + c4_state_blocks_per_row: tl.constexpr, + HAS_C4: tl.constexpr, + HAS_C128: tl.constexpr, BLOCK: tl.constexpr, ): pid = tl.program_id(0) - state_block = pid % blocks_per_row - job = pid // blocks_per_row - tail_row = job % 4 - layer = (job // 4) % layer_num - - layer_i64 = layer.to(tl.int64) - tail_row_i64 = tail_row.to(tl.int64) - tail_swa_slot = tl.load(resume_swa_slots + 2 * swa_page_size - 4 + tail_row_i64).to(tl.int64) - state_row = (tail_swa_slot // swa_page_size) * state_ring + tail_swa_slot % state_ring - cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) - - offsets = state_block * BLOCK + tl.arange(0, BLOCK) - offsets_i64 = offsets.to(tl.int64) - mask = offsets < state_width - source = ( - cpu_pages_f32 - + cpu_page * cpu_page_stride - + section_offset_f32 - + layer_i64 * section_layer_elems - + tail_row_i64 * state_width - + offsets_i64 - ) - target = state + layer_i64 * state_stride0 + state_row * state_stride1 + offsets_i64 - tl.store(target, tl.load(source, mask=mask), mask=mask) - - -def _launch_pack_pool_gpu_pages( - *, - full_slots, - mapping, - pool, - pool_page_size, - staging, - page_num, - token_page_size, - layer_num, - gpu_pages_per_cpu_page, - first_full_offset, - full_offset_per_gpu_page, - section_offset, - section_layer_nbytes, -): - gpu_page_nbytes = pool.shape[-1] - byte_blocks_per_gpu_page = triton.cdiv(gpu_page_nbytes, _BYTE_BLOCK) - _pack_pool_gpu_pages_kernel[(page_num * layer_num * gpu_pages_per_cpu_page * byte_blocks_per_gpu_page,)]( - full_slots, - mapping, - pool, - pool.stride(0), - pool.stride(1), - staging, - staging.stride(0), - token_page_size=token_page_size, - layer_num=layer_num, - gpu_pages_per_cpu_page=gpu_pages_per_cpu_page, - first_full_offset=first_full_offset, - full_offset_per_gpu_page=full_offset_per_gpu_page, - pool_page_size=pool_page_size, - gpu_page_nbytes=gpu_page_nbytes, - section_offset=section_offset, - section_layer_nbytes=section_layer_nbytes, - byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, - BLOCK=_BYTE_BLOCK, - num_warps=4, - ) - - -def _pack_c4_states(mem_manager, full_slots, state, staging, page_num, section_offset): - layout = mem_manager.cpu_cache_layout - state_width = state.shape[-1] - blocks_per_row = triton.cdiv(state_width, _STATE_BLOCK) - _pack_c4_tail_states_kernel[(page_num * mem_manager.n_c4 * 4 * blocks_per_row,)]( - full_slots, - mem_manager.full_to_swa_indexs, - state, - state.stride(0), - state.stride(1), - staging.view(torch.float32), - staging.stride(0) // 4, - token_page_size=layout.token_page_size, - layer_num=mem_manager.n_c4, - swa_page_size=mem_manager.swa_pool.page_size, - state_ring=mem_manager.c4_state_ring, - state_width=state_width, - section_offset_f32=section_offset // 4, - section_layer_elems=4 * state_width, - blocks_per_row=blocks_per_row, - BLOCK=_STATE_BLOCK, - num_warps=4, - ) + offsets_base = tl.arange(0, BLOCK) + + if HAS_C4: + if pid < c4_program_num: + byte_block = pid % c4_blocks_per_gpu_page + job = pid // c4_blocks_per_gpu_page + layer = job % c4_layer_num + history_block = job // c4_layer_num + history_block_i64 = history_block.to(tl.int64) + layer_i64 = layer.to(tl.int64) + absolute_block = history_block_i64 + first_history_block + cpu_page_list_index = absolute_block // blocks_per_cpu_page + block_in_cpu_page = absolute_block % blocks_per_cpu_page + cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) + pool_slot = tl.load(history_c4_slots + history_block_i64 * c4_pool_page_size).to(tl.int64) + physical_page = pool_slot // c4_pool_page_size + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + + c4_mask = offsets < c4_gpu_page_nbytes + c4_source = ( + cpu_pages + + cpu_page * cpu_stride0 + + c4_section_offset + + layer_i64 * c4_section_layer_nbytes + + block_in_cpu_page * c4_gpu_page_nbytes + + offsets_i64 + ) + c4_target = c4_pool + layer_i64 * c4_pool_stride0 + physical_page * c4_pool_stride1 + offsets_i64 + tl.store(c4_target, tl.load(c4_source, mask=c4_mask), mask=c4_mask) + + indexer_mask = offsets < c4_indexer_gpu_page_nbytes + indexer_source = ( + cpu_pages + + cpu_page * cpu_stride0 + + c4_indexer_section_offset + + layer_i64 * c4_indexer_section_layer_nbytes + + block_in_cpu_page * c4_indexer_gpu_page_nbytes + + offsets_i64 + ) + indexer_target = ( + c4_indexer_pool + + layer_i64 * c4_indexer_pool_stride0 + + physical_page * c4_indexer_pool_stride1 + + offsets_i64 + ) + tl.store(indexer_target, tl.load(indexer_source, mask=indexer_mask), mask=indexer_mask) + + c128_pid = pid - c4_program_num + if HAS_C128: + if (c128_pid >= 0) & (c128_pid < c128_program_num): + row = c128_pid % 2 + job = c128_pid // 2 + layer = job % c128_layer_num + history_block = job // c128_layer_num + history_block_i64 = history_block.to(tl.int64) + layer_i64 = layer.to(tl.int64) + row_i64 = row.to(tl.int64) + absolute_block = history_block_i64 + first_history_block + cpu_page_list_index = absolute_block // blocks_per_cpu_page + block_in_cpu_page = absolute_block % blocks_per_cpu_page + cpu_page = tl.load(cpu_page_indexes + cpu_page_list_index).to(tl.int64) + pool_slot = tl.load(history_c128_slots + history_block_i64 * 2 + row_i64).to(tl.int64) + physical_page = pool_slot // c128_pool_page_size + token_in_page = pool_slot % c128_pool_page_size + offsets_i64 = offsets_base.to(tl.int64) + cpu_row = block_in_cpu_page * 2 + row_i64 + source = ( + cpu_pages + + cpu_page * cpu_stride0 + + c128_section_offset + + layer_i64 * c128_section_layer_nbytes + + cpu_row * c128_row_nbytes + + offsets_i64 + ) + pool_page = c128_pool + layer_i64 * c128_pool_stride0 + physical_page * c128_pool_stride1 + data_mask = offsets_base < c128_data_nbytes + data_target = pool_page + token_in_page * c128_data_nbytes + offsets_i64 + tl.store(data_target, tl.load(source, mask=data_mask), mask=data_mask) + + scale_local = offsets_i64 - c128_data_nbytes + scale_mask = (offsets_base >= c128_data_nbytes) & (offsets_base < c128_row_nbytes) + scale_target = pool_page + c128_scale_offset + token_in_page * c128_scale_nbytes + scale_local + tl.store(scale_target, tl.load(source, mask=scale_mask), mask=scale_mask) + + swa_pid = c128_pid - c128_program_num + if (swa_pid >= 0) & (swa_pid < swa_program_num): + byte_block = swa_pid % swa_blocks_per_gpu_page + job = swa_pid // swa_blocks_per_gpu_page + gpu_page = job % swa_gpu_page_num + layer = job // swa_gpu_page_num + layer_i64 = layer.to(tl.int64) + gpu_page_i64 = gpu_page.to(tl.int64) + pool_slot = tl.load(resume_swa_slots + gpu_page_i64 * swa_pool_page_size).to(tl.int64) + physical_page = pool_slot // swa_pool_page_size + cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + mask = offsets < swa_gpu_page_nbytes + source = ( + cpu_pages + + cpu_page * cpu_stride0 + + swa_section_offset + + layer_i64 * swa_section_layer_nbytes + + gpu_page_i64 * swa_gpu_page_nbytes + + offsets_i64 + ) + target = swa_pool + layer_i64 * swa_pool_stride0 + physical_page * swa_pool_stride1 + offsets_i64 + tl.store(target, tl.load(source, mask=mask), mask=mask) + + state_pid = swa_pid - swa_program_num + if HAS_C4: + if state_pid >= 0: + byte_block = state_pid % c4_state_blocks_per_row + job = state_pid // c4_state_blocks_per_row + row = job % 4 + layer = job // 4 + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + swa_slot = tl.load(resume_swa_slots + 2 * swa_pool_page_size - 4 + row_i64).to(tl.int64) + state_row = (swa_slot // swa_pool_page_size) * c4_state_ring + swa_slot % c4_state_ring + cpu_page = tl.load(cpu_page_indexes + cpu_page_num - 1).to(tl.int64) + offsets = byte_block * BLOCK + offsets_base + offsets_i64 = offsets.to(tl.int64) + + state_mask = offsets < c4_state_row_nbytes + state_source = ( + cpu_pages + + cpu_page * cpu_stride0 + + c4_state_section_offset + + (layer_i64 * 4 + row_i64) * c4_state_row_nbytes + + offsets_i64 + ) + state_target = c4_state + layer_i64 * c4_state_stride0 + state_row * c4_state_stride1 + offsets_i64 + tl.store(state_target, tl.load(state_source, mask=state_mask), mask=state_mask) + + indexer_mask = offsets < c4_indexer_state_row_nbytes + indexer_source = ( + cpu_pages + + cpu_page * cpu_stride0 + + c4_indexer_state_section_offset + + (layer_i64 * 4 + row_i64) * c4_indexer_state_row_nbytes + + offsets_i64 + ) + indexer_target = ( + c4_indexer_state + + layer_i64 * c4_indexer_state_stride0 + + state_row * c4_indexer_state_stride1 + + offsets_i64 + ) + tl.store(indexer_target, tl.load(indexer_source, mask=indexer_mask), mask=indexer_mask) def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, staging: torch.Tensor) -> None: @@ -472,87 +491,97 @@ def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, sta assert staging.shape == (page_num, layout.page_nbytes) full_slots = source_mem_indexes.reshape(-1) - if mem_manager.n_c4: - for pool, section_offset, section_layer_nbytes in ( - (mem_manager.c4_pool, layout.c4_offset, layout.c4_layer_nbytes), - ( - mem_manager.c4_indexer_pool, - layout.c4_indexer_offset, - layout.c4_indexer_layer_nbytes, - ), - ): - _launch_pack_pool_gpu_pages( - full_slots=full_slots, - mapping=mem_manager.full_to_c4_indexs, - pool=pool.buffer, - pool_page_size=pool.page_size, - staging=staging, - page_num=page_num, - token_page_size=layout.token_page_size, - layer_num=mem_manager.n_c4, - gpu_pages_per_cpu_page=layout.c4_gpu_pages_per_cpu_page, - first_full_offset=3, - full_offset_per_gpu_page=_HISTORY_BLOCK_SIZE, - section_offset=section_offset, - section_layer_nbytes=section_layer_nbytes, - ) - - if mem_manager.n_c128: - pool = mem_manager.c128_pool - _pack_c128_rows_kernel[(page_num * mem_manager.n_c128 * layout.c128_rows_per_page,)]( - full_slots, - mem_manager.full_to_c128_indexs, - pool.buffer, - pool.buffer.stride(0), - pool.buffer.stride(1), - staging, - staging.stride(0), - token_page_size=layout.token_page_size, - layer_num=mem_manager.n_c128, - rows_per_page=layout.c128_rows_per_page, - pool_page_size=pool.page_size, - data_nbytes=pool.data_bytes_per_token, - scale_nbytes=pool.scale_bytes_per_token, - scale_offset=pool.scale_offset_in_page, - row_nbytes=layout.c128_row_nbytes, - section_offset=layout.c128_offset, - section_layer_nbytes=layout.c128_layer_nbytes, - BLOCK=triton.next_power_of_2(layout.c128_row_nbytes), - num_warps=4, - ) + has_c4 = mem_manager.c4_pool is not None + has_c128 = mem_manager.c128_pool is not None + c4_pool = mem_manager.c4_pool.buffer if has_c4 else None + c4_indexer_pool = mem_manager.c4_indexer_pool.buffer if has_c4 else None + c128_pool = mem_manager.c128_pool.buffer if has_c128 else None + swa_pool = mem_manager.swa_pool.buffer + c4_state = mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None + c4_indexer_state = mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None + + c4_blocks_per_gpu_page = ( + triton.cdiv(max(layout.c4_gpu_page_nbytes, layout.c4_indexer_gpu_page_nbytes), _BYTE_BLOCK) if has_c4 else 0 + ) + c4_program_num = mem_manager.n_c4 * layout.c4_gpu_pages_per_page * c4_blocks_per_gpu_page + c128_program_num = mem_manager.n_c128 * layout.c128_rows_per_page + swa_blocks_per_gpu_page = triton.cdiv(layout.swa_gpu_page_nbytes, _BYTE_BLOCK) + swa_program_num = mem_manager.layer_num * layout.swa_gpu_pages_per_page * swa_blocks_per_gpu_page + c4_state_blocks_per_row = ( + triton.cdiv(max(layout.c4_state_row_nbytes, layout.c4_indexer_state_row_nbytes), _BYTE_BLOCK) if has_c4 else 0 + ) + c4_state_program_num = mem_manager.n_c4 * layout.c4_state_rows * c4_state_blocks_per_row + programs_per_page = c4_program_num + c128_program_num + swa_program_num + c4_state_program_num - _launch_pack_pool_gpu_pages( - full_slots=full_slots, - mapping=mem_manager.full_to_swa_indexs, - pool=mem_manager.swa_pool.buffer, - pool_page_size=mem_manager.swa_pool.page_size, - staging=staging, - page_num=page_num, + _pack_gpu_cache_to_staging_kernel[(page_num * programs_per_page,)]( + full_slots, + mem_manager.full_to_c4_indexs if has_c4 else None, + c4_pool, + c4_pool.stride(0) if has_c4 else 0, + c4_pool.stride(1) if has_c4 else 0, + c4_indexer_pool, + c4_indexer_pool.stride(0) if has_c4 else 0, + c4_indexer_pool.stride(1) if has_c4 else 0, + mem_manager.full_to_c128_indexs if has_c128 else None, + c128_pool, + c128_pool.stride(0) if has_c128 else 0, + c128_pool.stride(1) if has_c128 else 0, + mem_manager.full_to_swa_indexs, + swa_pool, + swa_pool.stride(0), + swa_pool.stride(1), + c4_state, + c4_state.stride(0) if has_c4 else 0, + c4_state.stride(1) if has_c4 else 0, + c4_indexer_state, + c4_indexer_state.stride(0) if has_c4 else 0, + c4_indexer_state.stride(1) if has_c4 else 0, + staging, + staging.stride(0), + programs_per_page=programs_per_page, + c4_program_num=c4_program_num, + c128_program_num=c128_program_num, + swa_program_num=swa_program_num, token_page_size=layout.token_page_size, + history_block_size=_HISTORY_BLOCK_SIZE, + c128_ratio=_C128_RATIO, + c4_layer_num=mem_manager.n_c4, + c4_gpu_page_num=layout.c4_gpu_pages_per_page, + c4_pool_page_size=mem_manager.c4_pool.page_size if has_c4 else 0, + c4_gpu_page_nbytes=layout.c4_gpu_page_nbytes, + c4_section_offset=layout.c4_offset, + c4_section_layer_nbytes=layout.c4_layer_nbytes, + c4_indexer_gpu_page_nbytes=layout.c4_indexer_gpu_page_nbytes, + c4_indexer_section_offset=layout.c4_indexer_offset, + c4_indexer_section_layer_nbytes=layout.c4_indexer_layer_nbytes, + c4_blocks_per_gpu_page=c4_blocks_per_gpu_page, + c128_layer_num=mem_manager.n_c128, + c128_row_num=layout.c128_rows_per_page, + c128_pool_page_size=mem_manager.c128_pool.page_size if has_c128 else 0, + c128_data_nbytes=mem_manager.c128_pool.data_bytes_per_token if has_c128 else 0, + c128_scale_nbytes=mem_manager.c128_pool.scale_bytes_per_token if has_c128 else 0, + c128_scale_offset=mem_manager.c128_pool.scale_offset_in_page if has_c128 else 0, + c128_row_nbytes=layout.c128_row_nbytes, + c128_section_offset=layout.c128_offset, + c128_section_layer_nbytes=layout.c128_layer_nbytes, layer_num=mem_manager.layer_num, - gpu_pages_per_cpu_page=layout.swa_gpu_pages_per_cpu_page, - first_full_offset=layout.token_page_size - _HISTORY_BLOCK_SIZE, - full_offset_per_gpu_page=mem_manager.swa_pool.page_size, - section_offset=layout.swa_offset, - section_layer_nbytes=layout.swa_layer_nbytes, + swa_gpu_page_num=layout.swa_gpu_pages_per_page, + swa_pool_page_size=mem_manager.swa_pool.page_size, + swa_gpu_page_nbytes=layout.swa_gpu_page_nbytes, + swa_section_offset=layout.swa_offset, + swa_section_layer_nbytes=layout.swa_layer_nbytes, + swa_blocks_per_gpu_page=swa_blocks_per_gpu_page, + c4_state_ring=mem_manager.c4_state_ring, + c4_state_row_nbytes=layout.c4_state_row_nbytes, + c4_state_section_offset=layout.c4_state_offset, + c4_indexer_state_row_nbytes=layout.c4_indexer_state_row_nbytes, + c4_indexer_state_section_offset=layout.c4_indexer_state_offset, + c4_state_blocks_per_row=c4_state_blocks_per_row, + HAS_C4=has_c4, + HAS_C128=has_c128, + BLOCK=_BYTE_BLOCK, + num_warps=4, ) - if mem_manager.n_c4: - _pack_c4_states( - mem_manager, - full_slots, - mem_manager.c4_state_buffer, - staging, - page_num, - layout.c4_state_offset, - ) - _pack_c4_states( - mem_manager, - full_slots, - mem_manager.c4_indexer_state_buffer, - staging, - page_num, - layout.c4_indexer_state_offset, - ) def scatter_staging_to_cpu_pages( @@ -580,68 +609,6 @@ def scatter_staging_to_cpu_pages( ) -def _unpack_history_gpu_pages( - *, - pool_slots, - pool, - cpu_pages, - cpu_page_indexes, - first_history_block, - history_block_num, - blocks_per_cpu_page, - layer_num, - pool_slots_per_block, - section_offset, - section_layer_nbytes, -): - gpu_page_nbytes = pool.buffer.shape[-1] - byte_blocks_per_gpu_page = triton.cdiv(gpu_page_nbytes, _BYTE_BLOCK) - _unpack_history_gpu_pages_kernel[(history_block_num * layer_num * byte_blocks_per_gpu_page,)]( - pool_slots, - pool.buffer, - pool.buffer.stride(0), - pool.buffer.stride(1), - cpu_pages, - cpu_pages.stride(0), - cpu_page_indexes, - first_history_block, - layer_num=layer_num, - blocks_per_cpu_page=blocks_per_cpu_page, - pool_slots_per_block=pool_slots_per_block, - pool_page_size=pool.page_size, - gpu_page_nbytes=gpu_page_nbytes, - section_offset=section_offset, - section_layer_nbytes=section_layer_nbytes, - byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, - BLOCK=_BYTE_BLOCK, - num_warps=4, - ) - - -def _unpack_c4_states(mem_manager, resume_swa_slots, state, cpu_pages, page_indexes, section_offset): - state_width = state.shape[-1] - blocks_per_row = triton.cdiv(state_width, _STATE_BLOCK) - _unpack_c4_tail_states_kernel[(mem_manager.n_c4 * 4 * blocks_per_row,)]( - resume_swa_slots, - state, - state.stride(0), - state.stride(1), - cpu_pages.view(torch.float32), - cpu_pages.stride(0) // 4, - page_indexes, - page_indexes.numel(), - layer_num=mem_manager.n_c4, - swa_page_size=mem_manager.swa_pool.page_size, - state_ring=mem_manager.c4_state_ring, - state_width=state_width, - section_offset_f32=section_offset // 4, - section_layer_elems=4 * state_width, - blocks_per_row=blocks_per_row, - BLOCK=_STATE_BLOCK, - num_warps=4, - ) - - def unpack_cpu_cache_to_gpu( mem_manager, load_plan, @@ -674,92 +641,95 @@ def unpack_cpu_cache_to_gpu( expected_cpu_page_num = triton.cdiv(first_history_block + history_block_num, blocks_per_cpu_page) assert cpu_page_indexes.numel() == expected_cpu_page_num - if mem_manager.n_c4: + has_c4 = mem_manager.c4_pool is not None + has_c128 = mem_manager.c128_pool is not None + if has_c4: assert load_plan.history_c4_slots.shape == (history_block_num, mem_manager.c4_pool.page_size) - for pool, section_offset, section_layer_nbytes in ( - (mem_manager.c4_pool, layout.c4_offset, layout.c4_layer_nbytes), - ( - mem_manager.c4_indexer_pool, - layout.c4_indexer_offset, - layout.c4_indexer_layer_nbytes, - ), - ): - _unpack_history_gpu_pages( - pool_slots=load_plan.history_c4_slots, - pool=pool, - cpu_pages=cpu_pages, - cpu_page_indexes=cpu_page_indexes, - first_history_block=first_history_block, - history_block_num=history_block_num, - blocks_per_cpu_page=blocks_per_cpu_page, - layer_num=mem_manager.n_c4, - pool_slots_per_block=mem_manager.c4_pool.page_size, - section_offset=section_offset, - section_layer_nbytes=section_layer_nbytes, - ) - - if mem_manager.n_c128: + if has_c128: assert load_plan.history_c128_slots.shape == (history_block_num, 2) - pool = mem_manager.c128_pool - _unpack_c128_rows_kernel[(history_block_num * mem_manager.n_c128 * 2,)]( - load_plan.history_c128_slots, - pool.buffer, - pool.buffer.stride(0), - pool.buffer.stride(1), - cpu_pages, - cpu_pages.stride(0), - cpu_page_indexes, - first_history_block, - layer_num=mem_manager.n_c128, - blocks_per_cpu_page=blocks_per_cpu_page, - pool_page_size=pool.page_size, - data_nbytes=pool.data_bytes_per_token, - scale_nbytes=pool.scale_bytes_per_token, - scale_offset=pool.scale_offset_in_page, - row_nbytes=layout.c128_row_nbytes, - section_offset=layout.c128_offset, - section_layer_nbytes=layout.c128_layer_nbytes, - BLOCK=triton.next_power_of_2(layout.c128_row_nbytes), - num_warps=4, - ) - pool = mem_manager.swa_pool - byte_blocks_per_gpu_page = triton.cdiv(pool.buffer.shape[-1], _BYTE_BLOCK) - _unpack_resume_gpu_pages_kernel[ - (mem_manager.layer_num * layout.swa_gpu_pages_per_cpu_page * byte_blocks_per_gpu_page,) - ]( + c4_pool = mem_manager.c4_pool.buffer if has_c4 else None + c4_indexer_pool = mem_manager.c4_indexer_pool.buffer if has_c4 else None + c128_pool = mem_manager.c128_pool.buffer if has_c128 else None + swa_pool = mem_manager.swa_pool.buffer + c4_state = mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None + c4_indexer_state = mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None + + c4_blocks_per_gpu_page = ( + triton.cdiv(max(layout.c4_gpu_page_nbytes, layout.c4_indexer_gpu_page_nbytes), _BYTE_BLOCK) if has_c4 else 0 + ) + c4_program_num = history_block_num * mem_manager.n_c4 * c4_blocks_per_gpu_page + c128_program_num = history_block_num * mem_manager.n_c128 * 2 + swa_blocks_per_gpu_page = triton.cdiv(layout.swa_gpu_page_nbytes, _BYTE_BLOCK) + swa_program_num = mem_manager.layer_num * layout.swa_gpu_pages_per_page * swa_blocks_per_gpu_page + c4_state_blocks_per_row = ( + triton.cdiv(max(layout.c4_state_row_nbytes, layout.c4_indexer_state_row_nbytes), _BYTE_BLOCK) if has_c4 else 0 + ) + c4_state_program_num = mem_manager.n_c4 * layout.c4_state_rows * c4_state_blocks_per_row + + _unpack_cpu_cache_to_gpu_kernel[(c4_program_num + c128_program_num + swa_program_num + c4_state_program_num,)]( + load_plan.history_c4_slots if has_c4 else None, + load_plan.history_c128_slots if has_c128 else None, load_plan.resume_swa_slots, - pool.buffer, - pool.buffer.stride(0), - pool.buffer.stride(1), + c4_pool, + c4_pool.stride(0) if has_c4 else 0, + c4_pool.stride(1) if has_c4 else 0, + c4_indexer_pool, + c4_indexer_pool.stride(0) if has_c4 else 0, + c4_indexer_pool.stride(1) if has_c4 else 0, + c128_pool, + c128_pool.stride(0) if has_c128 else 0, + c128_pool.stride(1) if has_c128 else 0, + swa_pool, + swa_pool.stride(0), + swa_pool.stride(1), + c4_state, + c4_state.stride(0) if has_c4 else 0, + c4_state.stride(1) if has_c4 else 0, + c4_indexer_state, + c4_indexer_state.stride(0) if has_c4 else 0, + c4_indexer_state.stride(1) if has_c4 else 0, cpu_pages, cpu_pages.stride(0), cpu_page_indexes, cpu_page_indexes.numel(), + first_history_block, + c4_program_num, + c128_program_num, + c4_layer_num=mem_manager.n_c4, + c4_pool_page_size=mem_manager.c4_pool.page_size if has_c4 else 0, + c4_gpu_page_nbytes=layout.c4_gpu_page_nbytes, + c4_section_offset=layout.c4_offset, + c4_section_layer_nbytes=layout.c4_layer_nbytes, + c4_indexer_gpu_page_nbytes=layout.c4_indexer_gpu_page_nbytes, + c4_indexer_section_offset=layout.c4_indexer_offset, + c4_indexer_section_layer_nbytes=layout.c4_indexer_layer_nbytes, + c4_blocks_per_gpu_page=c4_blocks_per_gpu_page, + c128_layer_num=mem_manager.n_c128, + c128_pool_page_size=mem_manager.c128_pool.page_size if has_c128 else 0, + c128_data_nbytes=mem_manager.c128_pool.data_bytes_per_token if has_c128 else 0, + c128_scale_nbytes=mem_manager.c128_pool.scale_bytes_per_token if has_c128 else 0, + c128_scale_offset=mem_manager.c128_pool.scale_offset_in_page if has_c128 else 0, + c128_row_nbytes=layout.c128_row_nbytes, + c128_section_offset=layout.c128_offset, + c128_section_layer_nbytes=layout.c128_layer_nbytes, + blocks_per_cpu_page=blocks_per_cpu_page, layer_num=mem_manager.layer_num, - gpu_pages_per_cpu_page=layout.swa_gpu_pages_per_cpu_page, - pool_page_size=pool.page_size, - gpu_page_nbytes=pool.buffer.shape[-1], - section_offset=layout.swa_offset, - section_layer_nbytes=layout.swa_layer_nbytes, - byte_blocks_per_gpu_page=byte_blocks_per_gpu_page, + swa_program_num=swa_program_num, + swa_gpu_page_num=layout.swa_gpu_pages_per_page, + swa_pool_page_size=mem_manager.swa_pool.page_size, + swa_gpu_page_nbytes=layout.swa_gpu_page_nbytes, + swa_section_offset=layout.swa_offset, + swa_section_layer_nbytes=layout.swa_layer_nbytes, + swa_blocks_per_gpu_page=swa_blocks_per_gpu_page, + c4_state_ring=mem_manager.c4_state_ring, + c4_state_row_nbytes=layout.c4_state_row_nbytes, + c4_state_section_offset=layout.c4_state_offset, + c4_indexer_state_row_nbytes=layout.c4_indexer_state_row_nbytes, + c4_indexer_state_section_offset=layout.c4_indexer_state_offset, + c4_state_blocks_per_row=c4_state_blocks_per_row, + HAS_C4=has_c4, + HAS_C128=has_c128, BLOCK=_BYTE_BLOCK, num_warps=4, ) - if mem_manager.n_c4: - _unpack_c4_states( - mem_manager, - load_plan.resume_swa_slots, - mem_manager.c4_state_buffer, - cpu_pages, - cpu_page_indexes, - layout.c4_state_offset, - ) - _unpack_c4_states( - mem_manager, - load_plan.resume_swa_slots, - mem_manager.c4_indexer_state_buffer, - cpu_pages, - cpu_page_indexes, - layout.c4_indexer_state_offset, - ) diff --git a/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py new file mode 100644 index 0000000000..b26a3d81b2 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py @@ -0,0 +1,409 @@ +import torch +import triton +import triton.language as tl + +from .cache_staging_io import ( + BYTE_BLOCK as _BYTE_BLOCK, + STATE_BLOCK as _STATE_BLOCK, + copy_pool_pages, +) + +_C4_RATIO = 4 +_C4_POOL_PAGE_SIZE = 64 +_C4_TOKEN_BLOCK = _C4_RATIO * _C4_POOL_PAGE_SIZE +_C128_RATIO = 128 +_SWA_PAGE_SIZE = 128 + + +@triton.jit +def _pd_tail_kernel( + full_slots, + full_to_swa, + swa_pool, + swa_pool_stride0, + swa_pool_stride1, + c4_state, + c4_state_stride0, + c4_state_stride1, + c4_indexer_state, + c4_indexer_state_stride0, + c4_indexer_state_stride1, + c128_state, + c128_state_stride0, + c128_state_stride1, + staging, + staging_f32, + swa_program_num, + c4_program_num, + swa_page_num, + swa_first_full_offset, + swa_section_gpu_page_start, + c4_row_num, + c4_first_full_offset, + c4_section_row_start, + c128_row_num, + c128_first_position, + c128_section_row_start, + req_idx, + swa_page_size: tl.constexpr, + swa_page_nbytes: tl.constexpr, + swa_section_offset: tl.constexpr, + swa_section_layer_nbytes: tl.constexpr, + swa_blocks_per_page: tl.constexpr, + c4_state_ring: tl.constexpr, + c4_state_width: tl.constexpr, + c4_indexer_state_width: tl.constexpr, + c4_state_section_offset_f32: tl.constexpr, + c4_state_section_layer_elems: tl.constexpr, + c4_indexer_state_section_offset_f32: tl.constexpr, + c4_indexer_state_section_layer_elems: tl.constexpr, + c4_blocks_per_row: tl.constexpr, + c128_state_ring: tl.constexpr, + c128_state_width: tl.constexpr, + c128_state_section_offset_f32: tl.constexpr, + c128_state_section_layer_elems: tl.constexpr, + c128_blocks_per_row: tl.constexpr, + HAS_C4_STATE: tl.constexpr, + HAS_C128_STATE: tl.constexpr, + MODE: tl.constexpr, + BYTE_BLOCK: tl.constexpr, + STATE_BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + if pid < swa_program_num: + byte_block = pid % swa_blocks_per_page + job = pid // swa_blocks_per_page + gpu_page = job % swa_page_num + layer = job // swa_page_num + gpu_page_i64 = gpu_page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_slot = tl.load(full_slots + swa_first_full_offset + gpu_page_i64 * swa_page_size).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + physical_page = swa_slot // swa_page_size + offsets = byte_block * BYTE_BLOCK + tl.arange(0, BYTE_BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < swa_page_nbytes + pool_ptr = swa_pool + layer_i64 * swa_pool_stride0 + physical_page * swa_pool_stride1 + offsets_i64 + staging_ptr = ( + staging + + swa_section_offset + + layer_i64 * swa_section_layer_nbytes + + (swa_section_gpu_page_start + gpu_page_i64) * swa_page_nbytes + + offsets_i64 + ) + if MODE == 0: + tl.store(staging_ptr, tl.load(pool_ptr, mask=mask), mask=mask) + else: + tl.store(pool_ptr, tl.load(staging_ptr, mask=mask), mask=mask) + else: + state_pid = pid - swa_program_num + if HAS_C4_STATE: + if state_pid < c4_program_num: + state_block = state_pid % c4_blocks_per_row + job = state_pid // c4_blocks_per_row + row = job % c4_row_num + layer = job // c4_row_num + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_slot = tl.load(full_slots + c4_first_full_offset + row_i64).to(tl.int64) + swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + state_row = (swa_slot // swa_page_size) * c4_state_ring + swa_slot % c4_state_ring + offsets = state_block * STATE_BLOCK + tl.arange(0, STATE_BLOCK) + offsets_i64 = offsets.to(tl.int64) + c4_mask = offsets < c4_state_width + c4_state_ptr = c4_state + layer_i64 * c4_state_stride0 + state_row * c4_state_stride1 + offsets_i64 + c4_staging_ptr = ( + staging_f32 + + c4_state_section_offset_f32 + + layer_i64 * c4_state_section_layer_elems + + (c4_section_row_start + row_i64) * c4_state_width + + offsets_i64 + ) + indexer_mask = offsets < c4_indexer_state_width + indexer_state_ptr = ( + c4_indexer_state + + layer_i64 * c4_indexer_state_stride0 + + state_row * c4_indexer_state_stride1 + + offsets_i64 + ) + indexer_staging_ptr = ( + staging_f32 + + c4_indexer_state_section_offset_f32 + + layer_i64 * c4_indexer_state_section_layer_elems + + (c4_section_row_start + row_i64) * c4_indexer_state_width + + offsets_i64 + ) + if MODE == 0: + tl.store(c4_staging_ptr, tl.load(c4_state_ptr, mask=c4_mask), mask=c4_mask) + tl.store( + indexer_staging_ptr, + tl.load(indexer_state_ptr, mask=indexer_mask), + mask=indexer_mask, + ) + else: + tl.store(c4_state_ptr, tl.load(c4_staging_ptr, mask=c4_mask), mask=c4_mask) + tl.store( + indexer_state_ptr, + tl.load(indexer_staging_ptr, mask=indexer_mask), + mask=indexer_mask, + ) + + if HAS_C128_STATE: + c128_pid = state_pid - c4_program_num + if c128_pid >= 0: + state_block = c128_pid % c128_blocks_per_row + job = c128_pid // c128_blocks_per_row + row = job % c128_row_num + layer = job // c128_row_num + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + state_row = req_idx * c128_state_ring + (c128_first_position + row_i64) % c128_state_ring + offsets = state_block * STATE_BLOCK + tl.arange(0, STATE_BLOCK) + offsets_i64 = offsets.to(tl.int64) + mask = offsets < c128_state_width + state_ptr = c128_state + layer_i64 * c128_state_stride0 + state_row * c128_state_stride1 + offsets_i64 + staging_ptr = ( + staging_f32 + + c128_state_section_offset_f32 + + layer_i64 * c128_state_section_layer_elems + + (c128_section_row_start + row_i64) * c128_state_width + + offsets_i64 + ) + if MODE == 0: + tl.store(staging_ptr, tl.load(state_ptr, mask=mask), mask=mask) + else: + tl.store(state_ptr, tl.load(staging_ptr, mask=mask), mask=mask) + + +def _copy_pd_tail( + mode, + mem_manager, + layout, + full_slots, + staging, + start_kv_index, + end_kv_index, + request_kv_len, + req_idx, + swa_tail_start, +): + swa_intersection_start = max(start_kv_index, swa_tail_start) + swa_intersection_end = min(end_kv_index, request_kv_len) + swa_page_start = swa_intersection_start // _SWA_PAGE_SIZE * _SWA_PAGE_SIZE + swa_page_end = triton.cdiv(swa_intersection_end, _SWA_PAGE_SIZE) * _SWA_PAGE_SIZE + swa_page_num = (swa_page_end - swa_page_start) // _SWA_PAGE_SIZE + + c4_state = None + c4_indexer_state = None + c4_row_num = 0 + c4_first_full_offset = 0 + c4_section_row_start = 0 + if mem_manager.c4_pool is not None: + c4_remainder = request_kv_len % _C4_RATIO + c4_required_rows = _C4_RATIO if c4_remainder == 0 else _C4_RATIO + c4_remainder + c4_state_start = max(0, request_kv_len - c4_required_rows) + c4_intersection_start = max(start_kv_index, c4_state_start) + c4_intersection_end = min(end_kv_index, request_kv_len) + if c4_intersection_start < c4_intersection_end: + c4_state = mem_manager.c4_state_buffer + c4_indexer_state = mem_manager.c4_indexer_state_buffer + c4_row_num = c4_intersection_end - c4_intersection_start + c4_first_full_offset = c4_intersection_start - start_kv_index + c4_section_row_start = c4_intersection_start - c4_state_start + + c128_state = None + c128_row_num = 0 + c128_first_position = 0 + c128_section_row_start = 0 + c128_partial_row_num = request_kv_len % _C128_RATIO + if mem_manager.c128_pool is not None and c128_partial_row_num: + c128_state_start = request_kv_len - c128_partial_row_num + c128_intersection_start = max(start_kv_index, c128_state_start) + c128_intersection_end = min(end_kv_index, request_kv_len) + if c128_intersection_start < c128_intersection_end: + c128_state = mem_manager.c128_state_buffer + c128_row_num = c128_intersection_end - c128_intersection_start + c128_first_position = c128_intersection_start + c128_section_row_start = c128_intersection_start - c128_state_start + + swa_pool = mem_manager.swa_pool.buffer + swa_blocks_per_page = triton.cdiv(swa_pool.shape[-1], _BYTE_BLOCK) + swa_program_num = mem_manager.layer_num * swa_page_num * swa_blocks_per_page + + c4_state_width = c4_state.shape[-1] if c4_state is not None else 0 + c4_indexer_state_width = c4_indexer_state.shape[-1] if c4_indexer_state is not None else 0 + c4_blocks_per_row = ( + triton.cdiv(max(c4_state_width, c4_indexer_state_width), _STATE_BLOCK) if c4_state is not None else 0 + ) + c4_program_num = c4_state.shape[0] * c4_row_num * c4_blocks_per_row if c4_state is not None else 0 + + c128_state_width = c128_state.shape[-1] if c128_state is not None else 0 + c128_blocks_per_row = triton.cdiv(c128_state_width, _STATE_BLOCK) if c128_state is not None else 0 + c128_program_num = c128_state.shape[0] * c128_row_num * c128_blocks_per_row if c128_state is not None else 0 + + _pd_tail_kernel[(swa_program_num + c4_program_num + c128_program_num,)]( + full_slots, + mem_manager.full_to_swa_indexs, + swa_pool, + swa_pool.stride(0), + swa_pool.stride(1), + c4_state, + c4_state.stride(0) if c4_state is not None else 0, + c4_state.stride(1) if c4_state is not None else 0, + c4_indexer_state, + c4_indexer_state.stride(0) if c4_indexer_state is not None else 0, + c4_indexer_state.stride(1) if c4_indexer_state is not None else 0, + c128_state, + c128_state.stride(0) if c128_state is not None else 0, + c128_state.stride(1) if c128_state is not None else 0, + staging, + staging.view(torch.float32), + swa_program_num, + c4_program_num, + swa_page_num, + swa_page_start - start_kv_index, + (swa_page_start - swa_tail_start) // _SWA_PAGE_SIZE, + c4_row_num, + c4_first_full_offset, + c4_section_row_start, + c128_row_num, + c128_first_position, + c128_section_row_start, + req_idx, + swa_page_size=mem_manager.swa_pool.page_size, + swa_page_nbytes=swa_pool.shape[-1], + swa_section_offset=layout.swa_offset, + swa_section_layer_nbytes=layout.swa_layer_nbytes, + swa_blocks_per_page=swa_blocks_per_page, + c4_state_ring=mem_manager.c4_state_ring, + c4_state_width=c4_state_width, + c4_indexer_state_width=c4_indexer_state_width, + c4_state_section_offset_f32=layout.c4_state_offset // 4, + c4_state_section_layer_elems=layout.c4_state_rows * c4_state_width, + c4_indexer_state_section_offset_f32=layout.c4_indexer_state_offset // 4, + c4_indexer_state_section_layer_elems=layout.c4_state_rows * c4_indexer_state_width, + c4_blocks_per_row=c4_blocks_per_row, + c128_state_ring=mem_manager.c128_state_ring, + c128_state_width=c128_state_width, + c128_state_section_offset_f32=layout.c128_state_offset // 4, + c128_state_section_layer_elems=layout.c128_state_layer_nbytes // 4, + c128_blocks_per_row=c128_blocks_per_row, + HAS_C4_STATE=c4_state is not None, + HAS_C128_STATE=c128_state is not None, + MODE=0 if mode == "pack" else 1, + BYTE_BLOCK=_BYTE_BLOCK, + STATE_BLOCK=_STATE_BLOCK, + num_warps=4, + ) + + +def _copy_pd_cache_page( + mode, + mem_manager, + layout, + mem_indexes: torch.Tensor, + staging: torch.Tensor, + start_kv_index: int, + request_kv_len: int, + req_idx: int, +) -> int: + full_slots = mem_indexes.reshape(-1) + end_kv_index = start_kv_index + full_slots.numel() + staging = staging.reshape(-1) + c4_entry_num = end_kv_index // _C4_RATIO - start_kv_index // _C4_RATIO + c4_page_num = triton.cdiv(c4_entry_num, _C4_POOL_PAGE_SIZE) if c4_entry_num else 0 + c128_row_num = end_kv_index // _C128_RATIO - start_kv_index // _C128_RATIO + c4_pool = mem_manager.c4_pool + c128_pool = mem_manager.c128_pool + c4_work = c4_pool is not None and c4_page_num > 0 + c128_work = c128_pool is not None and c128_row_num > 0 + if c4_work or c128_work: + c4_indexer_pool = mem_manager.c4_indexer_pool if c4_work else None + copy_pool_pages( + mode, + full_slots=full_slots, + mapping=mem_manager.full_to_c4_indexs if c4_work else None, + pool=c4_pool.buffer if c4_work else None, + staging=staging, + page_num=1, + gpu_page_num=c4_page_num if c4_work else 0, + layer_num=c4_pool.layer_num if c4_work else 0, + full_slots_page_stride=0, + staging_page_stride=0, + first_full_offset=_C4_RATIO - 1, + full_offset_per_gpu_page=_C4_TOKEN_BLOCK, + pool_page_size=c4_pool.page_size if c4_work else 0, + section_offset=layout.c4_offset, + section_layer_nbytes=layout.c4_layer_nbytes, + paired_pool=c4_indexer_pool.buffer if c4_work else None, + paired_section_offset=layout.c4_indexer_offset, + paired_section_layer_nbytes=layout.c4_indexer_layer_nbytes, + c128_mapping=mem_manager.full_to_c128_indexs if c128_work else None, + c128_pool=c128_pool if c128_work else None, + c128_row_num=c128_row_num, + c128_first_full_offset=_C128_RATIO - 1, + c128_full_offset_per_row=_C128_RATIO, + c128_row_nbytes=layout.c128_row_nbytes, + c128_section_offset=layout.c128_offset, + c128_section_layer_nbytes=layout.c128_layer_nbytes, + ) + + swa_tail_start = max(0, request_kv_len // _C4_TOKEN_BLOCK * _C4_TOKEN_BLOCK - _C4_TOKEN_BLOCK) + if end_kv_index <= swa_tail_start: + return layout.swa_offset + + _copy_pd_tail( + mode, + mem_manager, + layout, + full_slots, + staging, + start_kv_index, + end_kv_index, + request_kv_len, + req_idx, + swa_tail_start, + ) + return layout.page_nbytes + + +def pack_pd_cache_page( + mem_manager, + layout, + mem_indexes: torch.Tensor, + staging: torch.Tensor, + start_kv_index: int, + request_kv_len: int, + req_idx: int, +) -> int: + return _copy_pd_cache_page( + "pack", + mem_manager, + layout, + mem_indexes, + staging, + start_kv_index, + request_kv_len, + req_idx, + ) + + +def unpack_pd_cache_page( + mem_manager, + layout, + mem_indexes: torch.Tensor, + staging: torch.Tensor, + start_kv_index: int, + request_kv_len: int, + req_idx: int, +) -> None: + _copy_pd_cache_page( + "unpack", + mem_manager, + layout, + mem_indexes, + staging, + start_kv_index, + request_kv_len, + req_idx, + ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 5939fc57c4..2df80b8402 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -83,15 +83,15 @@ def make_argument_parser() -> argparse.ArgumentParser: parser.add_argument( "--pd_kv_page_num", type=int, - default=16, - help="pd mode, kv move page_num", + default=None, + help="pd mode, kv move page_num; defaults to 8 for DeepSeek-V4 and 16 otherwise.", ) parser.add_argument( "--pd_kv_page_size", type=int, - default=1024, - help="pd mode, kv page size.", + default=None, + help="pd mode, kv page size; defaults to 2048 for DeepSeek-V4 and 1024 otherwise.", ) parser.add_argument( diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 585f42539d..089e96d526 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -82,7 +82,12 @@ def normal_or_p_d_start(args): auto_set_max_req_total_len(args) set_unique_server_name(args) - if args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4" and args.llm_kv_type in (None, "None"): + model_type = get_model_type(args.model_dir) + if args.pd_kv_page_num is None: + args.pd_kv_page_num = 8 if model_type == "deepseek_v4" else 16 + if args.pd_kv_page_size is None: + args.pd_kv_page_size = 2048 if model_type == "deepseek_v4" else 1024 + if args.enable_cpu_cache and model_type == "deepseek_v4" and args.llm_kv_type in (None, "None"): args.llm_kv_type = "fp8kv_dsa" if args.enable_mps: @@ -93,6 +98,10 @@ def normal_or_p_d_start(args): if args.run_mode not in ["normal", "prefill", "decode", "visual_only"]: return + if args.run_mode in ("prefill", "decode") and model_type == "deepseek_v4": + if args.tp != args.dp: + raise ValueError("DeepSeek-V4 PD requires one TP rank per DP replica (--tp must equal --dp)") + # 通过模型的参数判断是否是多模态模型,包含哪几种模态, 并设置是否启动相应得模块 if args.disable_vision is None: if has_vision_module(args.model_dir): diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index cc3fc9b523..ed5f7955ce 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -175,8 +175,8 @@ class StartArgs: mtp_draft_model_dir: Optional[str] = field(default=None) mtp_step: int = field(default=0) 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) + pd_kv_page_num: Optional[int] = field(default=None) + pd_kv_page_size: Optional[int] = field(default=None) pd_node_id: int = field(default=-1) enable_cpu_cache: bool = field(default=False) cpu_cache_storage_size: float = field(default=2) diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 1d68f81a9e..224b4c15b1 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -126,7 +126,7 @@ class PDAgentMetadata: agent_metadata: bytes num_pages: int page_reg_desc: Optional[bytes] = None - page_xfer_handles: Optional[int] = None + page_xfer_handles: Optional[Dict[int, object]] = None @dataclass @@ -134,6 +134,7 @@ class PDChunckedTransTask: request_id: int start_kv_index: int end_kv_index: int + request_kv_len: int time_out_secs: int pd_master_node_id: int @@ -162,6 +163,7 @@ class PDChunckedTransTask: # transfer params src_page_index: Optional[int] = None dst_page_index: Optional[int] = None + transfer_nbytes: Optional[int] = None # xfer_handle xfer_handle: Optional[int] = None diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 242e1089e7..8bc3f76740 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -1,5 +1,7 @@ import random +import torch import torch.multiprocessing as mp +from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.server.pd_io_struct import PDChunckedTransTask, PDChunckedTransTaskGroup, PDAbortReq from lightllm.server.router.model_infer.mode_backend.chunked_prefill.impl import ChunkedPrefillBackend from typing import List, Tuple @@ -122,10 +124,18 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): if self.radix_cache is not None: self.radix_cache.free_radix_cache_to_get_enough_token(need_mem_size) - mem_indexes = self.model.req_manager.mem_manager.alloc(need_size=need_mem_size) - self.model.req_manager.req_to_token_indexs[ + req_manager = self.model.req_manager + mem_indexes = req_manager.mem_manager.alloc(need_size=need_mem_size) + req_manager.req_to_token_indexs[ req_obj.req_idx, req_obj.cur_kv_len : (req_obj.cur_kv_len + need_mem_size) ] = mem_indexes + if isinstance(req_manager, DeepseekV4ReqManager): + req_manager.prepare_pd_decode_cache( + req_idx=req_obj.req_idx, + ready_cache_len=req_obj.cur_kv_len, + input_len=input_len, + new_full_slots=mem_indexes, + ) while req_obj.pd_trans_kv_start_index < input_len: cur_page_size = min(page_size, input_len - req_obj.pd_trans_kv_start_index) @@ -155,6 +165,8 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): group=group, page_kind="linear_att_state", ) + if isinstance(req_manager, DeepseekV4ReqManager): + torch.cuda.current_stream().synchronize() else: assert req_obj.cur_kv_len == input_len - 1 @@ -190,17 +202,14 @@ def _create_pd_trans_task( # only self.is_master_in_dp will be used. self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size - if page_kind == "kv": - req_idx = None - elif page_kind == "linear_att_state": - req_idx = req_obj.req_idx - else: + if page_kind not in ("kv", "linear_att_state"): raise ValueError(f"unknown PD trans page kind {page_kind}") trans_task = PDChunckedTransTask( request_id=req_obj.req_id, start_kv_index=kv_start_index, end_kv_index=kv_end_index, + request_kv_len=req_obj.shm_req.input_len, time_out_secs=180, pd_master_node_id=req_obj.sampling_param.pd_master_node_id, prefill_dp_index=None, @@ -219,7 +228,7 @@ def _create_pd_trans_task( first_gen_token_id=None, first_gen_token_logprob=None, page_kind=page_kind, - req_idx=req_idx, + req_idx=req_obj.req_idx, ) group.task_list.append(trans_task) req_obj.pd_task_num += 1 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 036c6f162b..77526dbf0c 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -163,17 +163,21 @@ def __init__( return def _warmup(self): - for dp_index in range(self.args.dp // self.args.nnodes): - with torch.cuda.stream(stream=self.copy_cuda_stream): - cur_mem = self.mem_managers[self.device_id] + cur_mem = self.mem_managers[self.device_id] + with torch.cuda.stream(stream=self.copy_cuda_stream): + cur_mem.kv_move_buffer[0].zero_() + for dp_index in range(self.args.dp // self.args.nnodes): cur_mem.read_page_kv_move_buffer_to_mem( - mem_indexes=[0], + mem_indexes=[cur_mem.HOLD_TOKEN_MEMINDEX], page_index=0, dp_index=dp_index, mem_managers=self.mem_managers, dp_world_size=self.dp_world_size, + start_kv_index=0, + request_kv_len=1, + req_idx=cur_mem.req_to_token_indexs.shape[0] - 1, ) - torch.cuda.current_stream().synchronize() + torch.cuda.current_stream().synchronize() return @log_exception @@ -289,6 +293,7 @@ def accept_peer_task_loop( local_trans_task.prefill_agent_metadata = remote_trans_task.prefill_agent_metadata local_trans_task.prefill_num_pages = remote_trans_task.prefill_num_pages local_trans_task.prefill_page_reg_desc = remote_trans_task.prefill_page_reg_desc + local_trans_task.transfer_nbytes = remote_trans_task.transfer_nbytes self.request_page_task_queue.put(local_trans_task) logger.info(f"recv WRITE request from prefill: {remote_trans_task.to_str()}") else: @@ -390,6 +395,8 @@ def read_page_to_mems_loop(self): dp_index=trans_task.decode_dp_index, mem_managers=self.mem_managers, dp_world_size=self.dp_world_size, + start_kv_index=trans_task.start_kv_index, + request_kv_len=trans_task.request_kv_len, page_kind=trans_task.page_kind, req_idx=trans_task.req_idx, ) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nccl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nccl_kv_transporter.py index 2ed0335ca5..63067826fc 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nccl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nccl_kv_transporter.py @@ -21,6 +21,12 @@ logger = init_logger(__name__) +def _page_prefix(page_tensor: Tensor, transfer_nbytes: int) -> Tensor: + if transfer_nbytes == page_tensor.numel() * page_tensor.element_size(): + return page_tensor + return page_tensor.reshape(-1)[: transfer_nbytes // page_tensor.element_size()] + + @dataclass class NcclAgentMetadata: agent_name: str @@ -298,7 +304,10 @@ def __init__(self, transporter: NcclKVTransporter, peer_name: str): def send_page(self, trans_task: PDChunckedTransTask) -> _NcclXferHandle: assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None - page_tensor = self.transporter.kv_move_buffer[trans_task.src_page_index] + assert trans_task.transfer_nbytes is not None + page_tensor = _page_prefix( + self.transporter.kv_move_buffer[trans_task.src_page_index], trans_task.transfer_nbytes + ) comm = self._ensure_comm(is_server=True) stream = self._get_stream() @@ -308,7 +317,8 @@ def send_page(self, trans_task: PDChunckedTransTask) -> _NcclXferHandle: logger.info( f"NCCL send page posted request_id={trans_task.request_id} " - f"src_page={trans_task.src_page_index} dst_agent={self.peer_name}" + f"src_page={trans_task.src_page_index} dst_agent={self.peer_name} " + f"transfer_nbytes={trans_task.transfer_nbytes}" ) return _NcclXferHandle(peer_name=self.peer_name, event=event) @@ -356,7 +366,10 @@ def _recv_page_loop(self, recv_queue: "queue.Queue[Optional[PDChunckedTransTask] def _recv_page(self, trans_task: PDChunckedTransTask): try: - page_tensor = self.transporter.kv_move_buffer[trans_task.dst_page_index] + assert trans_task.transfer_nbytes is not None + page_tensor = _page_prefix( + self.transporter.kv_move_buffer[trans_task.dst_page_index], trans_task.transfer_nbytes + ) comm = self._ensure_comm(is_server=False) stream = self._get_stream() comm.recv(page_tensor, src=0, stream=stream) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index bd5e11f05d..55ff077665 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -62,16 +62,42 @@ def _register_kv_move_buffer(self, kv_move_buffer: Tensor): self.dtype_byte_size = kv_move_buffer.element_size() self.page_len = self.page_size * self.num_layers * self.kv_head_num * self.head_dims * self.dtype_byte_size self.page_reg_desc = self.nixl_agent.register_memory(kv_move_buffer) - self.page_local_xfer_handles = self._create_paged_xfer_handles(self.page_reg_desc, self.num_pages) + self.page_local_xfer_handles = { + self.page_len: self._create_paged_xfer_handles(self.page_reg_desc, self.num_pages, self.page_len) + } - def _create_paged_xfer_handles(self, reg_desc: "nixlBind.nixlRegDList", page_num: int, agent_name: str = ""): + def _create_paged_xfer_handles( + self, + reg_desc: "nixlBind.nixlRegDList", + page_num: int, + transfer_nbytes: int, + agent_name: str = "", + ): base_addr, _, device_id, _ = reg_desc[0] pages_data = [] for page_id in range(page_num): - pages_data.append((base_addr + page_id * self.page_len, self.page_len, device_id)) + pages_data.append((base_addr + page_id * self.page_len, transfer_nbytes, device_id)) descs = self.nixl_agent.get_xfer_descs(pages_data, "VRAM") return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") + def _get_local_page_xfer_handles(self, transfer_nbytes: int): + if transfer_nbytes not in self.page_local_xfer_handles: + self.page_local_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( + self.page_reg_desc, self.num_pages, transfer_nbytes + ) + return self.page_local_xfer_handles[transfer_nbytes] + + def _get_remote_page_xfer_handles(self, remote_agent: PDAgentMetadata, transfer_nbytes: int): + if transfer_nbytes not in remote_agent.page_xfer_handles: + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + remote_agent.page_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( + page_mem_desc, + remote_agent.num_pages, + transfer_nbytes, + agent_name=remote_agent.agent_name, + ) + return remote_agent.page_xfer_handles[transfer_nbytes] + def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): if remote_agent.agent_name in self.remote_agents: return @@ -87,10 +113,14 @@ def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - kv_page_xfer_handles = self._create_paged_xfer_handles( - page_mem_desc, remote_agent.num_pages, agent_name=peer_name - ) - remote_agent.page_xfer_handles = kv_page_xfer_handles + remote_agent.page_xfer_handles = { + self.page_len: self._create_paged_xfer_handles( + page_mem_desc, + remote_agent.num_pages, + self.page_len, + agent_name=peer_name, + ) + } logger.info( f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" @@ -106,7 +136,8 @@ def remove_remote_agent(self, peer_name: str): assert remote_agent.agent_name == peer_name self.nixl_agent.remove_remote_agent(remote_agent.agent_name) if remote_agent.page_xfer_handles is not None: - self.nixl_agent.release_dlist_handle(remote_agent.page_xfer_handles) + for handles in remote_agent.page_xfer_handles.values(): + self.nixl_agent.release_dlist_handle(handles) except BaseException as e: logger.error(f"remove remote agent {peer_name} failed") logger.exception(str(e)) @@ -249,9 +280,10 @@ def write_blocks_paged( self.connect_add_remote_agent(_remote_agent) assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None + assert trans_task.transfer_nbytes is not None remote_agent: PDAgentMetadata = self.remote_agents[decode_agent_name] - src_handle = self.page_local_xfer_handles - dst_handle = remote_agent.page_xfer_handles + src_handle = self._get_local_page_xfer_handles(trans_task.transfer_nbytes) + dst_handle = self._get_remote_page_xfer_handles(remote_agent, trans_task.transfer_nbytes) handle = self.nixl_agent.make_prepped_xfer( "WRITE", src_handle, @@ -281,7 +313,8 @@ def release_xfer_handle(self, handle): def shutdown(self): self.nixl_agent.deregister_memory(self.page_reg_desc) - self.nixl_agent.release_dlist_handle(self.page_local_xfer_handles) + for handles in self.page_local_xfer_handles.values(): + self.nixl_agent.release_dlist_handle(handles) agent_names = list(self.remote_agents.keys()) for agent_name in agent_names: self.remove_remote_agent(agent_name) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 2a501f509b..ede4a9c865 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -113,16 +113,15 @@ def _create_pd_trans_task( .cpu() .tolist() ) - req_idx = None elif page_kind == "linear_att_state": mem_indexes = [] - req_idx = req_obj.req_idx else: raise ValueError(f"unknown PD trans page kind {page_kind}") trans_task = PDChunckedTransTask( request_id=req_obj.req_id, start_kv_index=kv_start_index, end_kv_index=kv_end_index, + request_kv_len=req_obj.shm_req.input_len, time_out_secs=182, pd_master_node_id=req_obj.sampling_param.pd_master_node_id, prefill_dp_index=self.dp_rank_in_node, @@ -141,7 +140,7 @@ def _create_pd_trans_task( first_gen_token_id=None, first_gen_token_logprob=None, page_kind=page_kind, - req_idx=req_idx, + req_idx=req_obj.req_idx, ) req_obj.pd_task_num += 1 return trans_task diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index e286a10f96..553eaf39b0 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -142,11 +142,14 @@ def _warmup(self): with torch.cuda.stream(stream=self.copy_cuda_stream): cur_mem = self.mem_managers[self.device_id] cur_mem.write_mem_to_page_kv_move_buffer( - mem_indexes=[0], + mem_indexes=[cur_mem.HOLD_TOKEN_MEMINDEX], page_index=0, dp_index=dp_index, mem_managers=self.mem_managers, dp_world_size=self.dp_world_size, + start_kv_index=0, + request_kv_len=1, + req_idx=cur_mem.req_to_token_indexs.shape[0] - 1, ) torch.cuda.current_stream().synchronize() return @@ -188,12 +191,14 @@ def local_copy_kv_loop(self): # 将kv 数据拷贝到 page 上,然后传输给 decode node,让其进行读取。 with torch.cuda.stream(stream=self.copy_cuda_stream): cur_mem = self.mem_managers[self.device_id] - cur_mem.write_mem_to_page_kv_move_buffer( + trans_task.transfer_nbytes = cur_mem.write_mem_to_page_kv_move_buffer( trans_task.mem_indexes, page_index=trans_task.src_page_index, dp_index=trans_task.prefill_dp_index, mem_managers=self.mem_managers, dp_world_size=self.dp_world_size, + start_kv_index=trans_task.start_kv_index, + request_kv_len=trans_task.request_kv_len, page_kind=trans_task.page_kind, req_idx=trans_task.req_idx, ) From 9fee60d51697ea134e5e80930acd5eb822db6e2f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 5 Aug 2026 14:27:42 +0000 Subject: [PATCH 099/214] support enable_dp_prompt_cache_fetch --- .../deepseek4_mem_manager.py | 6 + .../deepseek_v4/triton_kernel/dp_cache_io.py | 347 ++++++++++++++++++ .../dp_backend/dp_shared_kv_trans.py | 120 ++++-- 3 files changed, 446 insertions(+), 27 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 6494692b3e..6e0cb51e79 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -790,6 +790,12 @@ def register_swa_free_hook(self, fn) -> None: self._free_radix_unreferenced_swa_fn = fn return + def __getstate__(self): + state = self.__dict__.copy() + # The radix tree is process-local; IPC readers only need its CUDA cache tensors. + state["_free_radix_unreferenced_swa_fn"] = None + return state + def _alloc_swa_pages(self, need_pages: int) -> torch.Tensor: if need_pages > self.swa_page_allocator.can_use_mem_size and self._free_radix_unreferenced_swa_fn is not None: self._free_radix_unreferenced_swa_fn(need_pages - self.swa_page_allocator.can_use_mem_size) diff --git a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py new file mode 100644 index 0000000000..4b6b80a5d5 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py @@ -0,0 +1,347 @@ +import torch +import triton +import triton.language as tl + +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_PROMPT_CACHE_PAGE_SIZE + + +_C4_RATIO = 4 +_C128_RATIO = 128 +_BYTE_BLOCK = 8192 + + +@triton.jit +def _copy_dsv4_dp_cache_kernel( + src_full_slots, + dst_full_slots, + src_full_to_c4, + dst_full_to_c4, + src_c4_pool, + src_c4_pool_stride0, + src_c4_pool_stride1, + dst_c4_pool, + dst_c4_pool_stride0, + dst_c4_pool_stride1, + src_c4_indexer_pool, + src_c4_indexer_pool_stride0, + src_c4_indexer_pool_stride1, + dst_c4_indexer_pool, + dst_c4_indexer_pool_stride0, + dst_c4_indexer_pool_stride1, + src_full_to_c128, + dst_full_to_c128, + src_c128_pool, + src_c128_pool_stride0, + src_c128_pool_stride1, + dst_c128_pool, + dst_c128_pool_stride0, + dst_c128_pool_stride1, + src_full_to_swa, + dst_full_to_swa, + src_swa_pool, + src_swa_pool_stride0, + src_swa_pool_stride1, + dst_swa_pool, + dst_swa_pool_stride0, + dst_swa_pool_stride1, + src_c4_state, + src_c4_state_stride0, + src_c4_state_stride1, + dst_c4_state, + dst_c4_state_stride0, + dst_c4_state_stride1, + src_c4_indexer_state, + src_c4_indexer_state_stride0, + src_c4_indexer_state_stride1, + dst_c4_indexer_state, + dst_c4_indexer_state_stride0, + dst_c4_indexer_state_stride1, + token_num, + history_program_num, + history_block_size: tl.constexpr, + history_layer_num: tl.constexpr, + c4_ratio: tl.constexpr, + c4_pool_page_size: tl.constexpr, + c4_layer_num: tl.constexpr, + c4_pool_page_nbytes: tl.constexpr, + c4_indexer_pool_page_nbytes: tl.constexpr, + c4_copy_nbytes: tl.constexpr, + c128_ratio: tl.constexpr, + c128_layer_num: tl.constexpr, + c128_pool_page_size: tl.constexpr, + c128_data_nbytes: tl.constexpr, + c128_scale_nbytes: tl.constexpr, + c128_scale_offset: tl.constexpr, + swa_pool_page_size: tl.constexpr, + swa_pool_page_nbytes: tl.constexpr, + swa_program_num: tl.constexpr, + src_c4_state_ring: tl.constexpr, + dst_c4_state_ring: tl.constexpr, + c4_state_row_nbytes: tl.constexpr, + c4_indexer_state_row_nbytes: tl.constexpr, + c4_state_copy_nbytes: tl.constexpr, + HAS_HISTORY: tl.constexpr, + HAS_C4: tl.constexpr, + HAS_C128: tl.constexpr, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + lanes = tl.arange(0, BLOCK) + + if HAS_HISTORY: + if pid < history_program_num: + history_block = pid // history_layer_num + layer = pid % history_layer_num + history_block_i64 = history_block.to(tl.int64) + layer_i64 = layer.to(tl.int64) + + if HAS_C4: + if layer < c4_layer_num: + full_offset = history_block_i64 * history_block_size + c4_ratio - 1 + src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) + dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) + src_pool_slot = tl.load(src_full_to_c4 + src_full_slot).to(tl.int64) + dst_pool_slot = tl.load(dst_full_to_c4 + dst_full_slot).to(tl.int64) + src_page = src_pool_slot // c4_pool_page_size + dst_page = dst_pool_slot // c4_pool_page_size + + src_c4_page = src_c4_pool + layer_i64 * src_c4_pool_stride0 + src_page * src_c4_pool_stride1 + dst_c4_page = dst_c4_pool + layer_i64 * dst_c4_pool_stride0 + dst_page * dst_c4_pool_stride1 + src_indexer_page = ( + src_c4_indexer_pool + + layer_i64 * src_c4_indexer_pool_stride0 + + src_page * src_c4_indexer_pool_stride1 + ) + dst_indexer_page = ( + dst_c4_indexer_pool + + layer_i64 * dst_c4_indexer_pool_stride0 + + dst_page * dst_c4_indexer_pool_stride1 + ) + for byte_start in tl.range(0, c4_copy_nbytes, BLOCK): + offsets = byte_start + lanes + offsets_i64 = offsets.to(tl.int64) + c4_mask = offsets < c4_pool_page_nbytes + tl.store( + dst_c4_page + offsets_i64, + tl.load(src_c4_page + offsets_i64, mask=c4_mask), + mask=c4_mask, + ) + indexer_mask = offsets < c4_indexer_pool_page_nbytes + tl.store( + dst_indexer_page + offsets_i64, + tl.load(src_indexer_page + offsets_i64, mask=indexer_mask), + mask=indexer_mask, + ) + + if HAS_C128: + if layer < c128_layer_num: + for row in tl.static_range(0, 2): + full_offset = history_block_i64 * history_block_size + (row + 1) * c128_ratio - 1 + src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) + dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) + src_pool_slot = tl.load(src_full_to_c128 + src_full_slot).to(tl.int64) + dst_pool_slot = tl.load(dst_full_to_c128 + dst_full_slot).to(tl.int64) + src_page = src_pool_slot // c128_pool_page_size + dst_page = dst_pool_slot // c128_pool_page_size + src_token = src_pool_slot % c128_pool_page_size + dst_token = dst_pool_slot % c128_pool_page_size + + src_page_ptr = ( + src_c128_pool + layer_i64 * src_c128_pool_stride0 + src_page * src_c128_pool_stride1 + ) + dst_page_ptr = ( + dst_c128_pool + layer_i64 * dst_c128_pool_stride0 + dst_page * dst_c128_pool_stride1 + ) + offsets_i64 = lanes.to(tl.int64) + data_mask = lanes < c128_data_nbytes + src_data = src_page_ptr + src_token * c128_data_nbytes + offsets_i64 + dst_data = dst_page_ptr + dst_token * c128_data_nbytes + offsets_i64 + tl.store(dst_data, tl.load(src_data, mask=data_mask), mask=data_mask) + + scale_mask = lanes < c128_scale_nbytes + src_scale = src_page_ptr + c128_scale_offset + src_token * c128_scale_nbytes + offsets_i64 + dst_scale = dst_page_ptr + c128_scale_offset + dst_token * c128_scale_nbytes + offsets_i64 + tl.store(dst_scale, tl.load(src_scale, mask=scale_mask), mask=scale_mask) + + swa_pid = pid - history_program_num + if (swa_pid >= 0) & (swa_pid < swa_program_num): + page = swa_pid % 2 + layer = swa_pid // 2 + page_i64 = page.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_offset = token_num - history_block_size + page_i64 * swa_pool_page_size + src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) + dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) + src_pool_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) + dst_pool_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) + src_page = src_pool_slot // swa_pool_page_size + dst_page = dst_pool_slot // swa_pool_page_size + src_page_ptr = src_swa_pool + layer_i64 * src_swa_pool_stride0 + src_page * src_swa_pool_stride1 + dst_page_ptr = dst_swa_pool + layer_i64 * dst_swa_pool_stride0 + dst_page * dst_swa_pool_stride1 + for byte_start in tl.range(0, swa_pool_page_nbytes, BLOCK): + offsets = byte_start + lanes + offsets_i64 = offsets.to(tl.int64) + mask = offsets < swa_pool_page_nbytes + tl.store(dst_page_ptr + offsets_i64, tl.load(src_page_ptr + offsets_i64, mask=mask), mask=mask) + + if HAS_C4: + state_pid = swa_pid - swa_program_num + if state_pid >= 0: + row = state_pid % 4 + layer = state_pid // 4 + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + full_offset = token_num - 4 + row_i64 + src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) + dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) + src_swa_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) + dst_swa_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) + src_state_row = (src_swa_slot // swa_pool_page_size) * src_c4_state_ring + src_swa_slot % src_c4_state_ring + dst_state_row = (dst_swa_slot // swa_pool_page_size) * dst_c4_state_ring + dst_swa_slot % dst_c4_state_ring + + src_state_row_ptr = src_c4_state + layer_i64 * src_c4_state_stride0 + src_state_row * src_c4_state_stride1 + dst_state_row_ptr = dst_c4_state + layer_i64 * dst_c4_state_stride0 + dst_state_row * dst_c4_state_stride1 + src_indexer_row_ptr = ( + src_c4_indexer_state + + layer_i64 * src_c4_indexer_state_stride0 + + src_state_row * src_c4_indexer_state_stride1 + ) + dst_indexer_row_ptr = ( + dst_c4_indexer_state + + layer_i64 * dst_c4_indexer_state_stride0 + + dst_state_row * dst_c4_indexer_state_stride1 + ) + for byte_start in tl.range(0, c4_state_copy_nbytes, BLOCK): + offsets = byte_start + lanes + offsets_i64 = offsets.to(tl.int64) + state_mask = offsets < c4_state_row_nbytes + tl.store( + dst_state_row_ptr + offsets_i64, + tl.load(src_state_row_ptr + offsets_i64, mask=state_mask), + mask=state_mask, + ) + indexer_mask = offsets < c4_indexer_state_row_nbytes + tl.store( + dst_indexer_row_ptr + offsets_i64, + tl.load(src_indexer_row_ptr + offsets_i64, mask=indexer_mask), + mask=indexer_mask, + ) + + +def copy_dsv4_dp_cache( + src_mem_manager, + dst_mem_manager, + src_full_slots: torch.Tensor, + dst_full_slots: torch.Tensor, +) -> None: + """Copy one 256-token-aligned DP radix suffix directly between DSV4 GPU pools.""" + src_full_slots = src_full_slots.reshape(-1) + dst_full_slots = dst_full_slots.reshape(-1) + token_num = src_full_slots.numel() + assert ( + token_num == dst_full_slots.numel() + and token_num >= DSV4_PROMPT_CACHE_PAGE_SIZE + and token_num % DSV4_PROMPT_CACHE_PAGE_SIZE == 0 + ) + + has_c4 = dst_mem_manager.c4_pool is not None + has_c128 = dst_mem_manager.c128_pool is not None + history_layer_num = max(dst_mem_manager.n_c4, dst_mem_manager.n_c128) + history_program_num = token_num // DSV4_PROMPT_CACHE_PAGE_SIZE * history_layer_num + swa_program_num = dst_mem_manager.layer_num * 2 + c4_state_program_num = dst_mem_manager.n_c4 * 4 + + src_c4_pool = src_mem_manager.c4_pool.buffer if has_c4 else None + dst_c4_pool = dst_mem_manager.c4_pool.buffer if has_c4 else None + src_c4_indexer_pool = src_mem_manager.c4_indexer_pool.buffer if has_c4 else None + dst_c4_indexer_pool = dst_mem_manager.c4_indexer_pool.buffer if has_c4 else None + src_c128_pool = src_mem_manager.c128_pool.buffer if has_c128 else None + dst_c128_pool = dst_mem_manager.c128_pool.buffer if has_c128 else None + src_swa_pool = src_mem_manager.swa_pool.buffer + dst_swa_pool = dst_mem_manager.swa_pool.buffer + src_c4_state = src_mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None + dst_c4_state = dst_mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None + src_c4_indexer_state = src_mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None + dst_c4_indexer_state = dst_mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None + + c4_page_nbytes = dst_mem_manager.c4_pool.bytes_per_page if has_c4 else 0 + c4_indexer_page_nbytes = dst_mem_manager.c4_indexer_pool.bytes_per_page if has_c4 else 0 + c4_state_row_nbytes = dst_c4_state.shape[-1] if has_c4 else 0 + c4_indexer_state_row_nbytes = dst_c4_indexer_state.shape[-1] if has_c4 else 0 + + _copy_dsv4_dp_cache_kernel[(history_program_num + swa_program_num + c4_state_program_num,)]( + src_full_slots, + dst_full_slots, + src_mem_manager.full_to_c4_indexs if has_c4 else None, + dst_mem_manager.full_to_c4_indexs if has_c4 else None, + src_c4_pool, + src_c4_pool.stride(0) if has_c4 else 0, + src_c4_pool.stride(1) if has_c4 else 0, + dst_c4_pool, + dst_c4_pool.stride(0) if has_c4 else 0, + dst_c4_pool.stride(1) if has_c4 else 0, + src_c4_indexer_pool, + src_c4_indexer_pool.stride(0) if has_c4 else 0, + src_c4_indexer_pool.stride(1) if has_c4 else 0, + dst_c4_indexer_pool, + dst_c4_indexer_pool.stride(0) if has_c4 else 0, + dst_c4_indexer_pool.stride(1) if has_c4 else 0, + src_mem_manager.full_to_c128_indexs if has_c128 else None, + dst_mem_manager.full_to_c128_indexs if has_c128 else None, + src_c128_pool, + src_c128_pool.stride(0) if has_c128 else 0, + src_c128_pool.stride(1) if has_c128 else 0, + dst_c128_pool, + dst_c128_pool.stride(0) if has_c128 else 0, + dst_c128_pool.stride(1) if has_c128 else 0, + src_mem_manager.full_to_swa_indexs, + dst_mem_manager.full_to_swa_indexs, + src_swa_pool, + src_swa_pool.stride(0), + src_swa_pool.stride(1), + dst_swa_pool, + dst_swa_pool.stride(0), + dst_swa_pool.stride(1), + src_c4_state, + src_c4_state.stride(0) if has_c4 else 0, + src_c4_state.stride(1) if has_c4 else 0, + dst_c4_state, + dst_c4_state.stride(0) if has_c4 else 0, + dst_c4_state.stride(1) if has_c4 else 0, + src_c4_indexer_state, + src_c4_indexer_state.stride(0) if has_c4 else 0, + src_c4_indexer_state.stride(1) if has_c4 else 0, + dst_c4_indexer_state, + dst_c4_indexer_state.stride(0) if has_c4 else 0, + dst_c4_indexer_state.stride(1) if has_c4 else 0, + token_num, + history_program_num, + history_block_size=DSV4_PROMPT_CACHE_PAGE_SIZE, + history_layer_num=history_layer_num, + c4_ratio=_C4_RATIO, + c4_pool_page_size=dst_mem_manager.c4_pool.page_size if has_c4 else 0, + c4_layer_num=dst_mem_manager.n_c4, + c4_pool_page_nbytes=c4_page_nbytes, + c4_indexer_pool_page_nbytes=c4_indexer_page_nbytes, + c4_copy_nbytes=max(c4_page_nbytes, c4_indexer_page_nbytes), + c128_ratio=_C128_RATIO, + c128_layer_num=dst_mem_manager.n_c128, + c128_pool_page_size=dst_mem_manager.c128_pool.page_size if has_c128 else 0, + c128_data_nbytes=dst_mem_manager.c128_pool.data_bytes_per_token if has_c128 else 0, + c128_scale_nbytes=dst_mem_manager.c128_pool.scale_bytes_per_token if has_c128 else 0, + c128_scale_offset=dst_mem_manager.c128_pool.scale_offset_in_page if has_c128 else 0, + swa_pool_page_size=dst_mem_manager.swa_pool.page_size, + swa_pool_page_nbytes=dst_mem_manager.swa_pool.bytes_per_page, + swa_program_num=swa_program_num, + src_c4_state_ring=src_mem_manager.c4_state_ring, + dst_c4_state_ring=dst_mem_manager.c4_state_ring, + c4_state_row_nbytes=c4_state_row_nbytes, + c4_indexer_state_row_nbytes=c4_indexer_state_row_nbytes, + c4_state_copy_nbytes=max(c4_state_row_nbytes, c4_indexer_state_row_nbytes), + HAS_HISTORY=history_layer_num > 0, + HAS_C4=has_c4, + HAS_C128=has_c128, + BLOCK=_BYTE_BLOCK, + num_warps=4, + num_stages=1, + ) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index 2fa2c9cb9a..f982f4ede5 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -12,6 +12,7 @@ from lightllm.utils.dist_utils import get_current_device_id from lightllm.server.router.model_infer.infer_batch import g_infer_context import torch.distributed as dist +from lightllm.models.deepseek_v4.triton_kernel.dp_cache_io import copy_dsv4_dp_cache class DPKVSharedMoudle: @@ -34,11 +35,17 @@ def __init__(self, max_req_num: int, dp_size_in_node: int, backend): self.dp_rank_in_node = get_dp_rank_in_node() assert get_env_start_args().diverse_mode is False + if self.backend.is_deepseek_v4: + from lightllm.utils.device_utils import kv_trans_use_p2p + + assert kv_trans_use_p2p(), "DeepSeek-V4 DP prompt-cache fetch requires P2P KV transfer" + def fill_reqs_info(self, reqs: List[InferReq]): """ 填充请求的 kv 信息到共享内存中 """ - dist.barrier(group=self.backend.node_nccl_group) + if not self.backend.is_deepseek_v4: + dist.barrier(group=self.backend.node_nccl_group) if self.backend.is_master_in_dp: self.shared_req_infos.arr[0 : len(reqs), self.dp_rank_in_node, self._KV_LEN_INDEX] = [ req.cur_kv_len for req in reqs @@ -59,6 +66,15 @@ def build_shared_kv_trans_tasks( dist.barrier(group=self.backend.node_nccl_group) trans_tasks: List[TransTask] = [] + if self.backend.is_deepseek_v4: + dsv4_mem_manager = self.backend.model.mem_manager + ( + dsv4_swa_capacity, + dsv4_c4_capacity, + dsv4_c128_capacity, + ) = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() + dsv4_prompt_page_size = self.backend.model.req_manager.get_prompt_cache_page_size() + rank_max_radix_cache_lens = np.max( self.shared_req_infos.arr[0 : len(reqs), :, self._KV_LEN_INDEX], axis=1, keepdims=False ) @@ -71,9 +87,29 @@ def build_shared_kv_trans_tasks( # 计算需要传输的 kv 长度, 不能超过 req.get_cur_total_len() - 1 trans_size = min(max_req_radix_cache_len, req.get_cur_total_len() - 1) - req.cur_kv_len - if is_current_dp_handle and trans_size > 0 and g_infer_context.get_can_alloc_token_num() > trans_size: + can_alloc_dsv4_cache = True + if self.backend.is_deepseek_v4 and trans_size > 0: + need_swa_pages = dsv4_prompt_page_size // dsv4_mem_manager.swa_pool.page_size + need_c4_pages = trans_size // dsv4_prompt_page_size if dsv4_mem_manager.c4_pool is not None else 0 + need_c128_slots = trans_size // 128 if dsv4_mem_manager.c128_pool is not None else 0 + can_alloc_dsv4_cache = ( + dsv4_swa_capacity >= need_swa_pages + and dsv4_c4_capacity >= need_c4_pages + and dsv4_c128_capacity >= need_c128_slots + ) + + if ( + is_current_dp_handle + and trans_size > 0 + and g_infer_context.get_can_alloc_token_num() > trans_size + and can_alloc_dsv4_cache + ): g_infer_context.radix_cache.free_radix_cache_to_get_enough_token(trans_size) mem_indexes = self.backend.model.mem_manager.alloc(trans_size) + if self.backend.is_deepseek_v4: + dsv4_swa_capacity -= need_swa_pages + dsv4_c4_capacity -= need_c4_pages + dsv4_c128_capacity -= need_c128_slots max_kv_len_dp_rank = self.shared_req_infos.arr[req_index, :, self._KV_LEN_INDEX].argmax() max_kv_len_req_idx = int(self.shared_req_infos.arr[req_index, max_kv_len_dp_rank, self._REQ_IDX_INDEX]) max_kv_len_mem_manager_index = max_kv_len_dp_rank * self.backend.dp_world_size + self.backend.rank_in_dp @@ -98,33 +134,63 @@ def kv_trans(self, trans_tasks: List["TransTask"]): # kv 传输 if len(trans_tasks) > 0: - max_kv_len_mem_indexes = [] - max_kv_len_dp_ranks = [] - mem_indexes = [] - - for i, trans_task in enumerate(trans_tasks): - max_kv_len_mem_indexes.append(trans_task.max_kv_len_mem_indexes) - max_kv_len_dp_ranks.extend([trans_task.max_kv_len_dp_rank] * len(trans_task.max_kv_len_mem_indexes)) - mem_indexes.append(trans_task.mem_indexes) - - max_kv_len_mem_indexes_tensor = torch.cat(max_kv_len_mem_indexes).to(dtype=torch.int64, device="cuda") - max_kv_len_dp_ranks_tensor = torch.tensor(max_kv_len_dp_ranks, dtype=torch.int32, device="cuda") - mem_indexes_tensor = torch.cat(mem_indexes).to(dtype=torch.int64, device="cuda") - self.backend.model.mem_manager.operator.copy_kv_from_other_dp_ranks( - mem_managers=self.backend.mem_managers, - move_token_indexes=max_kv_len_mem_indexes_tensor, - token_dp_indexes=max_kv_len_dp_ranks_tensor, - mem_indexes=mem_indexes_tensor, - dp_size_in_node=self.backend.dp_size_in_node, - rank_in_dp=self.backend.rank_in_dp, - ) - self.backend.logger.info(f"dp_i {self.dp_rank_in_node} transfer kv tokens num: {len(mem_indexes_tensor)}") + if self.backend.is_deepseek_v4: + req_manager = g_infer_context.req_manager + for trans_task in trans_tasks: + start = trans_task.req.cur_kv_len + end = start + len(trans_task.mem_indexes) + dst_full_slots = req_manager.req_to_token_indexs[trans_task.req.req_idx, start:end] + dst_full_slots.copy_(trans_task.mem_indexes, non_blocking=True) + req_manager.prepare_pd_decode_cache( + req_idx=trans_task.req.req_idx, + ready_cache_len=start, + input_len=end, + new_full_slots=dst_full_slots, + ) + dst_mem_manager = self.backend.model.mem_manager + dst_req_to_token = self.backend.model.req_manager.req_to_token_indexs + for trans_task in trans_tasks: + src_mem_manager = self.backend.mem_managers[trans_task.max_kv_len_mem_manager_index] + start = trans_task.req.cur_kv_len + end = start + len(trans_task.mem_indexes) + src_mem_indexes = trans_task.max_kv_len_mem_indexes + dst_mem_indexes = dst_req_to_token[trans_task.req.req_idx, start:end] + copy_dsv4_dp_cache(src_mem_manager, dst_mem_manager, src_mem_indexes, dst_mem_indexes) + else: + max_kv_len_mem_indexes = [] + max_kv_len_dp_ranks = [] + mem_indexes = [] + + for i, trans_task in enumerate(trans_tasks): + max_kv_len_mem_indexes.append(trans_task.max_kv_len_mem_indexes) + max_kv_len_dp_ranks.extend([trans_task.max_kv_len_dp_rank] * len(trans_task.max_kv_len_mem_indexes)) + mem_indexes.append(trans_task.mem_indexes) + + max_kv_len_mem_indexes_tensor = torch.cat(max_kv_len_mem_indexes).to(dtype=torch.int64, device="cuda") + max_kv_len_dp_ranks_tensor = torch.tensor(max_kv_len_dp_ranks, dtype=torch.int32, device="cuda") + mem_indexes_tensor = torch.cat(mem_indexes).to(dtype=torch.int64, device="cuda") + self.backend.model.mem_manager.operator.copy_kv_from_other_dp_ranks( + mem_managers=self.backend.mem_managers, + move_token_indexes=max_kv_len_mem_indexes_tensor, + token_dp_indexes=max_kv_len_dp_ranks_tensor, + mem_indexes=mem_indexes_tensor, + dp_size_in_node=self.backend.dp_size_in_node, + rank_in_dp=self.backend.rank_in_dp, + ) + + transfer_token_num = sum(len(trans_task.mem_indexes) for trans_task in trans_tasks) + self.backend.logger.info(f"dp_i {self.dp_rank_in_node} transfer kv tokens num: {transfer_token_num}") + + if self.backend.is_deepseek_v4: + # Source radix mappings stay alive until every peer copy on the current stream finishes. + dist.barrier(group=self.backend.node_nccl_group) for trans_task in trans_tasks: - g_infer_context.req_manager.req_to_token_indexs[ - trans_task.req.req_idx, - trans_task.req.cur_kv_len : (trans_task.req.cur_kv_len + len(trans_task.mem_indexes)), - ] = trans_task.mem_indexes + if not self.backend.is_deepseek_v4: + g_infer_context.req_manager.req_to_token_indexs[ + trans_task.req.req_idx, + trans_task.req.cur_kv_len : (trans_task.req.cur_kv_len + len(trans_task.mem_indexes)), + ] = trans_task.mem_indexes trans_task.req.cur_kv_len += len(trans_task.mem_indexes) if self.backend.is_master_in_dp: trans_task.req.shm_req.shm_cur_kv_len = trans_task.req.cur_kv_len From 97ae2d12acdd77bcf87a786a710c0f1f049d8196 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 5 Aug 2026 15:16:25 +0000 Subject: [PATCH 100/214] fix(deepep): derive NVSHMEM QP depth from decode capacity --- lightllm/utils/envs_utils.py | 21 ++++++++++++++++++++- 1 file changed, 20 insertions(+), 1 deletion(-) diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 4289123a4f..7c3f4184b0 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -37,6 +37,18 @@ def set_env_start_args(args): if not isinstance(args, dict): args = vars(args) os.environ["LIGHTLLM_START_ARGS"] = json.dumps(args) + if args["enable_ep_moe"]: + decode_capacity = get_deepep_num_max_dispatch_tokens_per_rank_decode() + min_qp_depth = 2 * (decode_capacity + 1) + derived_qp_depth = 1 << (min_qp_depth - 1).bit_length() + configured_qp_depth = int(os.getenv("NVSHMEM_QP_DEPTH", derived_qp_depth)) + if configured_qp_depth < derived_qp_depth: + logger.warning( + "NVSHMEM_QP_DEPTH=%d is below the required minimum; using %d instead.", + configured_qp_depth, + derived_qp_depth, + ) + os.environ["NVSHMEM_QP_DEPTH"] = str(max(configured_qp_depth, derived_qp_depth)) return @@ -86,7 +98,14 @@ def get_deepep_num_max_dispatch_tokens_per_rank_decode(): args = get_env_start_args() required = args.running_max_req_size * (args.mtp_step + 1) required = ((required + 7) // 8) * 8 - return int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", required)) + configured = int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", required)) + if configured != required: + logger.warning( + "NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE=%d differs from the automatically derived value %d.", + configured, + required, + ) + return max(configured, required) def get_lightllm_gunicorn_keep_alive(): From 4552d583030d12956744476c64254e7508cd5855 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 7 Aug 2026 01:51:12 +0000 Subject: [PATCH 101/214] handle consecutive DeepSeek DSML tool-call blocks --- lightllm/server/function_call_parser.py | 236 +++++++++++++----------- 1 file changed, 126 insertions(+), 110 deletions(-) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index 72d4901f93..8ee42ac345 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1513,6 +1513,8 @@ def __init__(self, block_name: str = "function_calls"): self._last_arguments = "" self._accumulated_params: List[tuple] = [] self._in_function_calls = False # Track if we're inside a function_calls block + # Text after a closed block is held unless it is whitespace before another block. + self._after_function_calls = False def has_tool_call(self, text: str) -> bool: return self.bot_token in text @@ -1532,143 +1534,157 @@ def _dsml_params_to_json(self, params: List[tuple]) -> str: def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult: """One-time parsing for DSML format tool calls.""" - block_start = text.find(self.bot_token) - if block_start == -1: + first_block_start = text.find(self.bot_token) + if first_block_start == -1: return StreamingParseResult(normal_text=text, calls=[]) - block_body_start = block_start + len(self.bot_token) - block_end = text.find(self.eot_token, block_body_start) - if block_end == -1: - return StreamingParseResult(normal_text=text, calls=[]) - - normal_text = text[:block_start].removesuffix("\n\n") - + normal_text = text[:first_block_start].removesuffix("\n\n") calls = [] + search_pos = first_block_start + + while True: + block_start = text.find(self.bot_token, search_pos) + if block_start == -1: + break + if text[search_pos:block_start].strip(): + break + + block_body_start = block_start + len(self.bot_token) + block_end = text.find(self.eot_token, block_body_start) + if block_end == -1: + break + + invoke_matches = self.invoke_regex.findall(text[block_body_start:block_end]) + for func_name, invoke_body in invoke_matches: + param_matches = self.param_regex.findall(invoke_body) + args_json = self._dsml_params_to_json(param_matches) + match_result = { + "name": func_name, + "parameters": json.loads(args_json), + } + for item in self.parse_base_json(match_result, tools): + item.tool_index = len(calls) + calls.append(item) - invoke_matches = self.invoke_regex.findall(text[block_body_start:block_end]) - for func_name, invoke_body in invoke_matches: - param_matches = self.param_regex.findall(invoke_body) - args_json = self._dsml_params_to_json(param_matches) - match_result = { - "name": func_name, - "parameters": json.loads(args_json), - } - for item in self.parse_base_json(match_result, tools): - item.tool_index = len(calls) - calls.append(item) + search_pos = block_end + len(self.eot_token) return StreamingParseResult(normal_text=normal_text, calls=calls) def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> StreamingParseResult: """Streaming incremental parsing for DSML format tool calls.""" self._buffer += new_text - current_text = self._buffer - - # Check if we're inside a function_calls block or starting one - has_tool = self.has_tool_call(current_text) or self._in_function_calls + normal_text_parts = [] + calls: List[ToolCallItem] = [] - if not has_tool: - partial_len = self._ends_with_partial_token(current_text, self.bot_token) - if partial_len: - return StreamingParseResult() + try: + while True: + current_text = self._buffer - normal_text = current_text - self._buffer = "" - for e_token in [self.eot_token, self.invoke_end_token]: - if e_token in normal_text: - normal_text = normal_text.replace(e_token, "") - return StreamingParseResult(normal_text=normal_text) + if not self._in_function_calls: + block_start = current_text.find(self.bot_token) + if block_start == -1: + if self._after_function_calls: + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - normal_text = "" + partial_len = self._ends_with_partial_token(current_text, self.bot_token) + if partial_len: + normal_text_parts.append(current_text[:-partial_len]) + self._buffer = current_text[-partial_len:] + else: + normal_text_parts.append(current_text) + self._buffer = "" + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + + outside_text = current_text[:block_start] + if self._after_function_calls: + if outside_text.strip(): + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + else: + normal_text_parts.append(outside_text.removesuffix("\n\n")) + + self._buffer = current_text[block_start + len(self.bot_token) :] + self._in_function_calls = True + self._after_function_calls = False + continue - # Mark that we're inside a function_calls block - if self.has_tool_call(current_text): - block_start = current_text.find(self.bot_token) - normal_text = current_text[:block_start].removesuffix("\n\n") - current_text = current_text[block_start:] - self._buffer = current_text - self._in_function_calls = True + self._buffer = current_text.lstrip() + current_text = self._buffer + if not current_text: + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - # Check if function_calls block has ended - if self.eot_token in current_text: - self._in_function_calls = False + if current_text.startswith(self.eot_token): + self._buffer = current_text[len(self.eot_token) :] + self._in_function_calls = False + self._after_function_calls = True + continue - calls: List[ToolCallItem] = [] + if self.eot_token.startswith(current_text): + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - try: - # Try to find complete invoke blocks first - while True: - complete_invoke_match = self.invoke_regex.search(current_text) - if not complete_invoke_match: - break - func_name = complete_invoke_match.group(1) - invoke_body = complete_invoke_match.group(2) + complete_invoke_match = self.invoke_regex.match(current_text) + if complete_invoke_match: + func_name = complete_invoke_match.group(1) + invoke_body = complete_invoke_match.group(2) - if self.current_tool_id == -1: - self.current_tool_id = 0 - self.prev_tool_call_arr = [] - self.streamed_args_for_tool = [""] - self._accumulated_params = [] + if self.current_tool_id == -1: + self.current_tool_id = 0 + self.prev_tool_call_arr = [] + self.streamed_args_for_tool = [""] + self._accumulated_params = [] - while len(self.prev_tool_call_arr) <= self.current_tool_id: - self.prev_tool_call_arr.append({}) - while len(self.streamed_args_for_tool) <= self.current_tool_id: - self.streamed_args_for_tool.append("") + while len(self.prev_tool_call_arr) <= self.current_tool_id: + self.prev_tool_call_arr.append({}) + while len(self.streamed_args_for_tool) <= self.current_tool_id: + self.streamed_args_for_tool.append("") - param_matches = self.param_regex.findall(invoke_body) - args_json = self._dsml_params_to_json(param_matches) + param_matches = self.param_regex.findall(invoke_body) + args_json = self._dsml_params_to_json(param_matches) - if not self.current_tool_name_sent: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=func_name, - parameters="", + if not self.current_tool_name_sent: + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + name=func_name, + parameters="", + ) ) - ) - self.current_tool_name_sent = True + self.current_tool_name_sent = True - # Send complete arguments (or remaining diff) - sent = len(self.streamed_args_for_tool[self.current_tool_id]) - argument_diff = args_json[sent:] - if argument_diff: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=None, - parameters=argument_diff, + sent = len(self.streamed_args_for_tool[self.current_tool_id]) + argument_diff = args_json[sent:] + if argument_diff: + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + name=None, + parameters=argument_diff, + ) ) - ) - self.streamed_args_for_tool[self.current_tool_id] += argument_diff + self.streamed_args_for_tool[self.current_tool_id] += argument_diff - try: - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": json.loads(args_json), - } - except json.JSONDecodeError: - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": {}, - } + try: + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": json.loads(args_json), + } + except json.JSONDecodeError: + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": {}, + } - # Remove processed invoke from buffer - invoke_end_pos = current_text.find(self.invoke_end_token, complete_invoke_match.start()) - if invoke_end_pos != -1: - self._buffer = current_text[invoke_end_pos + len(self.invoke_end_token) :] - else: self._buffer = current_text[complete_invoke_match.end() :] + self.current_tool_id += 1 + self._last_arguments = "" + self.current_tool_name_sent = False + self._accumulated_params = [] + self.streamed_args_for_tool.append("") + continue - self.current_tool_id += 1 - self._last_arguments = "" - self.current_tool_name_sent = False - self._accumulated_params = [] - self.streamed_args_for_tool.append("") - current_text = self._buffer + partial_match = self.partial_invoke_regex.match(current_text) + if not partial_match: + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - # Partial invoke: name is known but parameters are still streaming - partial_match = self.partial_invoke_regex.search(current_text) - if partial_match: func_name = partial_match.group(1) partial_body = partial_match.group(2) @@ -1722,11 +1738,11 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami except json.JSONDecodeError: pass - return StreamingParseResult(normal_text=normal_text, calls=calls) + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) except Exception as e: logger.error(f"Error in DeepSeekV32 parse_streaming_increment: {e}") - return StreamingParseResult(normal_text=normal_text, calls=calls) + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) class Qwen3CoderDetector(BaseFormatDetector): From f64e82eedd57b4811ff65655dfd51fafb173b2aa Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 7 Aug 2026 08:15:32 +0000 Subject: [PATCH 102/214] fix: support UE8M0 FP8 quantization on SM100 for TP --- .../quantization/fp8act_quant_kernel.py | 101 ++++++++++-------- .../fp8w8a8_block_quant_kernel.py | 33 ++++-- lightllm/common/quantization/deepgemm.py | 4 +- 3 files changed, 83 insertions(+), 55 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py index 0a68372887..a439a01f9d 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py @@ -14,6 +14,14 @@ pass +@triton.jit +def _ceil_to_ue8m0(x): + bits = x.to(tl.float32).to(tl.int32, bitcast=True) + exp = ((bits >> 23) & 0xFF) + ((bits & 0x7FFFFF) != 0) + exp = tl.maximum(tl.minimum(exp, 254), 1) + return (exp << 23).to(tl.float32, bitcast=True) + + # Adapted from https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/layers/quantization/fp8_kernel.py @triton.jit def _per_token_group_quant_fp8( @@ -25,21 +33,19 @@ def _per_token_group_quant_fp8( eps, fp8_min, fp8_max, - xs_m, xs_n, - xs_row_major: tl.constexpr, + xs_stride_m, + xs_stride_n, BLOCK: tl.constexpr, NEED_MASK: tl.constexpr, + USE_UE8M0_SCALE: tl.constexpr, ): g_id = tl.program_id(0) y_ptr += g_id * y_stride y_q_ptr += g_id * y_stride - if xs_row_major: - y_s_ptr += g_id - else: - row_id = g_id // xs_n - col_id = g_id % xs_n - y_s_ptr += col_id * xs_m + row_id # col major + row_id = g_id // xs_n + col_id = g_id % xs_n + y_s_ptr += row_id * xs_stride_m + col_id * xs_stride_n cols = tl.arange(0, BLOCK) # N <= BLOCK @@ -52,8 +58,11 @@ def _per_token_group_quant_fp8( y = tl.load(y_ptr + cols, mask=mask, other=other).to(tl.float32) # Quant - _absmax = tl.maximum(tl.max(tl.abs(y)), eps) - y_s = _absmax / fp8_max + _absmax = tl.max(tl.abs(y)) + if USE_UE8M0_SCALE: + y_s = _ceil_to_ue8m0(tl.maximum(_absmax, 1.0e-4) / fp8_max) + else: + y_s = tl.maximum(_absmax, eps) / fp8_max y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty) tl.store(y_q_ptr + cols, y_q, mask=mask) @@ -67,6 +76,7 @@ def lightllm_per_token_group_quant_fp8( x_s: torch.Tensor, eps: float = 1e-10, dtype: torch.dtype = torch.float8_e4m3fn, + use_ue8m0_scales: bool = False, ): """group-wise, per-token quantization on input tensor `x`. Args: @@ -80,8 +90,8 @@ def lightllm_per_token_group_quant_fp8( assert x.shape[-1] % group_size == 0, "the last dimension of `x` cannot be divisible by `group_size`" assert x.is_contiguous(), "`x` is not contiguous" - xs_row_major = x_s.is_contiguous() - xs_m, xs_n = x_s.shape + xs_n = x_s.shape[-1] + xs_stride_m, xs_stride_n = x_s.stride() finfo = torch.finfo(dtype) fp8_max = finfo.max @@ -102,11 +112,12 @@ def lightllm_per_token_group_quant_fp8( eps, fp8_min=fp8_min, fp8_max=fp8_max, - xs_m=xs_m, xs_n=xs_n, - xs_row_major=xs_row_major, + xs_stride_m=xs_stride_m, + xs_stride_n=xs_stride_n, BLOCK=BLOCK, NEED_MASK=BLOCK != group_size, + USE_UE8M0_SCALE=use_ue8m0_scales, num_warps=num_warps, num_stages=num_stages, ) @@ -121,50 +132,48 @@ def per_token_group_quant_fp8( column_major_scales: bool = False, scale_tma_aligned: bool = False, alloc_func: Callable = torch.empty, + use_ue8m0_scales: bool = False, ): x_q = alloc_func(x.shape, dtype=dtype, device=x.device) - x_s = None - # Adapted from - # https://github.com/sgl-project/sglang/blob/7e257cd666c0d639626487987ea8e590da1e9395/python/sglang/srt/layers/quantization/fp8_kernel.py#L290 - if HAS_SGL_KERNEL: - finfo = torch.finfo(dtype) - fp8_max, fp8_min = finfo.max, finfo.min - - # 创建scale张量 - if column_major_scales: - if scale_tma_aligned: - # 对齐到4 * sizeof(float) - aligned_size = (x.shape[-2] + 3) // 4 * 4 - x_s = alloc_func( - x.shape[:-2] + (x.shape[-1] // group_size, aligned_size), - device=x.device, - dtype=torch.float32, - ).permute(-1, -2)[: x.shape[-2], :] - else: - x_s = alloc_func( - (x.shape[-1] // group_size,) + x.shape[:-1], - device=x.device, - dtype=torch.float32, - ).permute(-1, -2) + if column_major_scales: + if scale_tma_aligned: + aligned_size = (x.shape[-2] + 3) // 4 * 4 + x_s = alloc_func( + x.shape[:-2] + (x.shape[-1] // group_size, aligned_size), + device=x.device, + dtype=torch.float32, + ).permute(-1, -2)[: x.shape[-2], :] else: x_s = alloc_func( - x.shape[:-1] + (x.shape[-1] // group_size,), + (x.shape[-1] // group_size,) + x.shape[:-1], device=x.device, dtype=torch.float32, - ) - - # 使用SGL kernel进行量化 - sgl_ops.sgl_per_token_group_quant_fp8(x, x_q, x_s, group_size, 1e-10, fp8_min, fp8_max, False, enable_v2=True) + ).permute(-1, -2) else: - # 使用LightLLM kernel进行量化 x_s = alloc_func( x.shape[:-1] + (x.shape[-1] // group_size,), device=x.device, dtype=torch.float32, ) - lightllm_per_token_group_quant_fp8(x, group_size, x_q, x_s, eps=1e-10, dtype=torch.float8_e4m3fn) - if column_major_scales and scale_tma_aligned: - x_s = tma_align_input_scale(x_s) + + # Adapted from + # https://github.com/sgl-project/sglang/blob/7e257cd666c0d639626487987ea8e590da1e9395/python/sglang/srt/layers/quantization/fp8_kernel.py#L290 + if HAS_SGL_KERNEL and not use_ue8m0_scales: + finfo = torch.finfo(dtype) + fp8_max, fp8_min = finfo.max, finfo.min + # 使用SGL kernel进行量化 + sgl_ops.sgl_per_token_group_quant_fp8(x, x_q, x_s, group_size, 1e-10, fp8_min, fp8_max, False, enable_v2=True) + else: + # 使用LightLLM kernel进行量化 + lightllm_per_token_group_quant_fp8( + x, + group_size, + x_q, + x_s, + eps=eps, + dtype=dtype, + use_ue8m0_scales=use_ue8m0_scales, + ) return x_q, x_s diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8w8a8_block_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8w8a8_block_quant_kernel.py index 3881cfe4b8..c711aae850 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8w8a8_block_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8w8a8_block_quant_kernel.py @@ -5,7 +5,15 @@ @triton.jit -def weight_quant_kernel(x_ptr, s_ptr, y_ptr, M, N, BLOCK_SIZE: tl.constexpr): +def _ceil_to_ue8m0(x): + bits = x.to(tl.float32).to(tl.int32, bitcast=True) + exp = ((bits >> 23) & 0xFF) + ((bits & 0x7FFFFF) != 0) + exp = tl.maximum(tl.minimum(exp, 254), 1) + return (exp << 23).to(tl.float32, bitcast=True) + + +@triton.jit +def weight_quant_kernel(x_ptr, s_ptr, y_ptr, M, N, BLOCK_SIZE: tl.constexpr, USE_UE8M0_SCALE: tl.constexpr): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) n_blocks = tl.cdiv(N, BLOCK_SIZE) @@ -20,14 +28,21 @@ def weight_quant_kernel(x_ptr, s_ptr, y_ptr, M, N, BLOCK_SIZE: tl.constexpr): amax = tl.max(tl.abs(x)) max_fp8e4m3_val = 448.0 - scale = amax / max_fp8e4m3_val - y = (x / (scale + 1e-6)).to(y_ptr.dtype.element_ty) + if USE_UE8M0_SCALE: + scale = _ceil_to_ue8m0(tl.maximum(amax, 1.0e-4) / max_fp8e4m3_val) + denom = scale + else: + scale = amax / max_fp8e4m3_val + denom = scale + 1e-6 + y = (x / denom).to(y_ptr.dtype.element_ty) tl.store(y_ptr + offs, y, mask=mask) tl.store(s_ptr + pid_m * n_blocks + pid_n, scale) -def mm_weight_quant(x: torch.Tensor, block_size: int = 128) -> tuple[torch.Tensor, torch.Tensor]: +def mm_weight_quant( + x: torch.Tensor, block_size: int = 128, use_ue8m0_scales: bool = False +) -> tuple[torch.Tensor, torch.Tensor]: assert x.is_contiguous(), "Input tensor must be contiguous" M, N = x.size() @@ -38,11 +53,13 @@ def mm_weight_quant(x: torch.Tensor, block_size: int = 128) -> tuple[torch.Tenso s_scales = torch.empty((num_blocks_m, num_blocks_n), dtype=torch.float32, device=x.device) grid = lambda meta: (triton.cdiv(M, meta["BLOCK_SIZE"]), triton.cdiv(N, meta["BLOCK_SIZE"])) - weight_quant_kernel[grid](x, s_scales, y_quant, M, N, BLOCK_SIZE=block_size) + weight_quant_kernel[grid](x, s_scales, y_quant, M, N, BLOCK_SIZE=block_size, USE_UE8M0_SCALE=use_ue8m0_scales) return y_quant, s_scales -def weight_quant(x: torch.Tensor, block_size: int = 128) -> tuple[torch.Tensor, torch.Tensor]: +def weight_quant( + x: torch.Tensor, block_size: int = 128, use_ue8m0_scales: bool = False +) -> tuple[torch.Tensor, torch.Tensor]: assert x.is_contiguous(), "Input tensor must be contiguous" x = x.cuda(get_current_device_id()) if x.dim() == 3: @@ -51,8 +68,8 @@ def weight_quant(x: torch.Tensor, block_size: int = 128) -> tuple[torch.Tensor, num_blocks_n = triton.cdiv(x.shape[2], block_size) s_scales = torch.empty((x.shape[0], num_blocks_m, num_blocks_n), dtype=torch.float32, device=x.device) for i in range(x.shape[0]): - y_quant[i], s_scales[i] = mm_weight_quant(x[i], block_size) + y_quant[i], s_scales[i] = mm_weight_quant(x[i], block_size, use_ue8m0_scales=use_ue8m0_scales) return y_quant, s_scales else: - y_quant, s_scales = mm_weight_quant(x, block_size) + y_quant, s_scales = mm_weight_quant(x, block_size, use_ue8m0_scales=use_ue8m0_scales) return y_quant, s_scales diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 7e7abafcae..1cb134a7e1 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -5,6 +5,7 @@ from lightllm.common.quantization.registry import QUANTMETHODS from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import per_token_group_quant_fp8 from lightllm.utils.log_utils import init_logger +from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -62,7 +63,7 @@ def quantize(self, weight: torch.Tensor, output: WeightPack): from lightllm.common.basemodel.triton_kernel.quantization.fp8w8a8_block_quant_kernel import weight_quant device = output.weight.device - weight, scale = weight_quant(weight.cuda(device), self.block_size) + weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=is_sm100_gpu()) output.weight.copy_(weight) output.weight_scale.copy_(scale) return @@ -90,6 +91,7 @@ def apply( column_major_scales=True, scale_tma_aligned=True, alloc_func=alloc_func, + use_ue8m0_scales=is_sm100_gpu(), ) if out is None: From edfb743f0cc46d55d071bed7357244b430390d43 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 10 Aug 2026 02:46:55 +0000 Subject: [PATCH 103/214] batch DP cache preparation and P2P transfers --- lightllm/common/req_manager.py | 44 ++- .../deepseek_v4/triton_kernel/dp_cache_io.py | 262 ++++++++---------- .../model_infer/mode_backend/base_backend.py | 2 + .../dp_backend/dp_shared_kv_trans.py | 97 +++++-- .../pd/decode_node_impl/decode_impl.py | 6 +- 5 files changed, 237 insertions(+), 174 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index c62c842779..a0dcdbcb8e 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -530,33 +530,47 @@ def prepare_prefill( def prepare_pd_decode_cache( self, - req_idx: int, - ready_cache_len: int, - input_len: int, + req_list: List[int], + ready_list: List[int], + seq_list: List[int], new_full_slots: torch.Tensor, ) -> None: - """Allocate DSV4 derived slots for a suffix received by a PD decode node.""" + """Allocate DSV4 derived slots for request-major suffixes received from peers.""" page = self.get_prompt_cache_page_size() - assert ready_cache_len % page == 0 - assert new_full_slots.numel() == input_len - ready_cache_len + assert len(req_list) == len(ready_list) == len(seq_list) and len(req_list) > 0 + assert all(ready % page == 0 and seq_len > ready for ready, seq_len in zip(ready_list, seq_list)) + assert new_full_slots.numel() == sum(seq_len - ready for ready, seq_len in zip(ready_list, seq_list)) new_full_slots = new_full_slots.reshape(-1).to(self.req_to_token_indexs.device, non_blocking=True) self.prepare_prefill_compress_slots( - req_list=[req_idx], - ready_list=[ready_cache_len], - seq_list=[input_len], + req_list=req_list, + ready_list=ready_list, + seq_list=seq_list, mem_indexes=new_full_slots, ) - resume_start = max(ready_cache_len, max(0, input_len // page * page - page)) + # swa 只保存最后一部分,前面的不需要 + swa_start_list = [] + swa_parts = [] + offset = 0 + # swa 不一样 + for ready, seq_len in zip(ready_list, seq_list): + swa_start = max(ready, max(0, seq_len // page * page - page)) + swa_start_list.append(swa_start) + suffix_len = seq_len - ready + swa_parts.append(new_full_slots[offset + swa_start - ready : offset + suffix_len]) + offset += suffix_len + swa_full_slots = swa_parts[0] if len(swa_parts) == 1 else torch.cat(swa_parts) + self.mem_manager.alloc_swa_prefill( - new_full_slots[resume_start - ready_cache_len :], + swa_full_slots, self.req_to_token_indexs, - req_list=[req_idx], - ready_list=[resume_start], - seq_list=[input_len], + req_list=req_list, + ready_list=swa_start_list, + seq_list=seq_list, ) - self._swa_evict_marks[req_idx] = resume_start + for req_idx, swa_start in zip(req_list, swa_start_list): + self._swa_evict_marks[req_idx] = swa_start return def prepare_decode_swa( diff --git a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py index 4b6b80a5d5..4b929235b8 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py @@ -8,56 +8,42 @@ _C4_RATIO = 4 _C128_RATIO = 128 _BYTE_BLOCK = 8192 +# Per-source row: c4 map/data/indexer, c128 map/data, SWA map/data, c4 state/indexer state. +_SOURCE_POOL_PTR_COUNT = 9 +# Per-task row: source manager index, token count, source full-slot pointer, destination full-slot pointer. +_TASK_META_WIDTH = 4 @triton.jit -def _copy_dsv4_dp_cache_kernel( - src_full_slots, - dst_full_slots, - src_full_to_c4, +def _copy_dsv4_dp_caches_kernel( + source_pool_ptrs, + task_meta, + history_meta, dst_full_to_c4, - src_c4_pool, - src_c4_pool_stride0, - src_c4_pool_stride1, dst_c4_pool, dst_c4_pool_stride0, dst_c4_pool_stride1, - src_c4_indexer_pool, - src_c4_indexer_pool_stride0, - src_c4_indexer_pool_stride1, dst_c4_indexer_pool, dst_c4_indexer_pool_stride0, dst_c4_indexer_pool_stride1, - src_full_to_c128, dst_full_to_c128, - src_c128_pool, - src_c128_pool_stride0, - src_c128_pool_stride1, dst_c128_pool, dst_c128_pool_stride0, dst_c128_pool_stride1, - src_full_to_swa, dst_full_to_swa, - src_swa_pool, - src_swa_pool_stride0, - src_swa_pool_stride1, dst_swa_pool, dst_swa_pool_stride0, dst_swa_pool_stride1, - src_c4_state, - src_c4_state_stride0, - src_c4_state_stride1, dst_c4_state, dst_c4_state_stride0, dst_c4_state_stride1, - src_c4_indexer_state, - src_c4_indexer_state_stride0, - src_c4_indexer_state_stride1, dst_c4_indexer_state, dst_c4_indexer_state_stride0, dst_c4_indexer_state_stride1, - token_num, history_program_num, + task_num, + source_pool_ptr_count: tl.constexpr, + task_meta_width: tl.constexpr, history_block_size: tl.constexpr, history_layer_num: tl.constexpr, c4_ratio: tl.constexpr, @@ -75,8 +61,7 @@ def _copy_dsv4_dp_cache_kernel( swa_pool_page_size: tl.constexpr, swa_pool_page_nbytes: tl.constexpr, swa_program_num: tl.constexpr, - src_c4_state_ring: tl.constexpr, - dst_c4_state_ring: tl.constexpr, + c4_state_ring: tl.constexpr, c4_state_row_nbytes: tl.constexpr, c4_indexer_state_row_nbytes: tl.constexpr, c4_state_copy_nbytes: tl.constexpr, @@ -90,14 +75,23 @@ def _copy_dsv4_dp_cache_kernel( if HAS_HISTORY: if pid < history_program_num: - history_block = pid // history_layer_num + history_index = pid // history_layer_num layer = pid % history_layer_num - history_block_i64 = history_block.to(tl.int64) + task = tl.load(history_meta + history_index * 2).to(tl.int64) + history_block = tl.load(history_meta + history_index * 2 + 1).to(tl.int64) + task_row = task_meta + task * task_meta_width + source_manager = tl.load(task_row).to(tl.int64) + src_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) + dst_full_slots = tl.load(task_row + 3).to(tl.pointer_type(tl.int32)) + source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count layer_i64 = layer.to(tl.int64) if HAS_C4: if layer < c4_layer_num: - full_offset = history_block_i64 * history_block_size + c4_ratio - 1 + src_full_to_c4 = tl.load(source_ptr_row).to(tl.pointer_type(tl.int32)) + src_c4_pool = tl.load(source_ptr_row + 1).to(tl.pointer_type(tl.uint8)) + src_c4_indexer_pool = tl.load(source_ptr_row + 2).to(tl.pointer_type(tl.uint8)) + full_offset = history_block * history_block_size + c4_ratio - 1 src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) src_pool_slot = tl.load(src_full_to_c4 + src_full_slot).to(tl.int64) @@ -105,12 +99,12 @@ def _copy_dsv4_dp_cache_kernel( src_page = src_pool_slot // c4_pool_page_size dst_page = dst_pool_slot // c4_pool_page_size - src_c4_page = src_c4_pool + layer_i64 * src_c4_pool_stride0 + src_page * src_c4_pool_stride1 + src_c4_page = src_c4_pool + layer_i64 * dst_c4_pool_stride0 + src_page * dst_c4_pool_stride1 dst_c4_page = dst_c4_pool + layer_i64 * dst_c4_pool_stride0 + dst_page * dst_c4_pool_stride1 src_indexer_page = ( src_c4_indexer_pool - + layer_i64 * src_c4_indexer_pool_stride0 - + src_page * src_c4_indexer_pool_stride1 + + layer_i64 * dst_c4_indexer_pool_stride0 + + src_page * dst_c4_indexer_pool_stride1 ) dst_indexer_page = ( dst_c4_indexer_pool @@ -135,8 +129,10 @@ def _copy_dsv4_dp_cache_kernel( if HAS_C128: if layer < c128_layer_num: + src_full_to_c128 = tl.load(source_ptr_row + 3).to(tl.pointer_type(tl.int32)) + src_c128_pool = tl.load(source_ptr_row + 4).to(tl.pointer_type(tl.uint8)) for row in tl.static_range(0, 2): - full_offset = history_block_i64 * history_block_size + (row + 1) * c128_ratio - 1 + full_offset = history_block * history_block_size + (row + 1) * c128_ratio - 1 src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) src_pool_slot = tl.load(src_full_to_c128 + src_full_slot).to(tl.int64) @@ -147,7 +143,7 @@ def _copy_dsv4_dp_cache_kernel( dst_token = dst_pool_slot % c128_pool_page_size src_page_ptr = ( - src_c128_pool + layer_i64 * src_c128_pool_stride0 + src_page * src_c128_pool_stride1 + src_c128_pool + layer_i64 * dst_c128_pool_stride0 + src_page * dst_c128_pool_stride1 ) dst_page_ptr = ( dst_c128_pool + layer_i64 * dst_c128_pool_stride0 + dst_page * dst_c128_pool_stride1 @@ -163,159 +159,146 @@ def _copy_dsv4_dp_cache_kernel( dst_scale = dst_page_ptr + c128_scale_offset + dst_token * c128_scale_nbytes + offsets_i64 tl.store(dst_scale, tl.load(src_scale, mask=scale_mask), mask=scale_mask) - swa_pid = pid - history_program_num - if (swa_pid >= 0) & (swa_pid < swa_program_num): - page = swa_pid % 2 - layer = swa_pid // 2 - page_i64 = page.to(tl.int64) - layer_i64 = layer.to(tl.int64) - full_offset = token_num - history_block_size + page_i64 * swa_pool_page_size - src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) - dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_pool_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) - dst_pool_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) - src_page = src_pool_slot // swa_pool_page_size - dst_page = dst_pool_slot // swa_pool_page_size - src_page_ptr = src_swa_pool + layer_i64 * src_swa_pool_stride0 + src_page * src_swa_pool_stride1 - dst_page_ptr = dst_swa_pool + layer_i64 * dst_swa_pool_stride0 + dst_page * dst_swa_pool_stride1 - for byte_start in tl.range(0, swa_pool_page_nbytes, BLOCK): - offsets = byte_start + lanes - offsets_i64 = offsets.to(tl.int64) - mask = offsets < swa_pool_page_nbytes - tl.store(dst_page_ptr + offsets_i64, tl.load(src_page_ptr + offsets_i64, mask=mask), mask=mask) + tail_pid = pid - history_program_num + if tail_pid >= 0: + task = (tail_pid % task_num).to(tl.int64) + task_pid = tail_pid // task_num + task_row = task_meta + task * task_meta_width + source_manager = tl.load(task_row).to(tl.int64) + token_num = tl.load(task_row + 1).to(tl.int64) + src_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) + dst_full_slots = tl.load(task_row + 3).to(tl.pointer_type(tl.int32)) + source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count + src_full_to_swa = tl.load(source_ptr_row + 5).to(tl.pointer_type(tl.int32)) + src_swa_pool = tl.load(source_ptr_row + 6).to(tl.pointer_type(tl.uint8)) - if HAS_C4: - state_pid = swa_pid - swa_program_num - if state_pid >= 0: - row = state_pid % 4 - layer = state_pid // 4 - row_i64 = row.to(tl.int64) + if task_pid < swa_program_num: + page = task_pid % 2 + layer = task_pid // 2 + page_i64 = page.to(tl.int64) layer_i64 = layer.to(tl.int64) - full_offset = token_num - 4 + row_i64 + full_offset = token_num - history_block_size + page_i64 * swa_pool_page_size src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_swa_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) - dst_swa_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) - src_state_row = (src_swa_slot // swa_pool_page_size) * src_c4_state_ring + src_swa_slot % src_c4_state_ring - dst_state_row = (dst_swa_slot // swa_pool_page_size) * dst_c4_state_ring + dst_swa_slot % dst_c4_state_ring - - src_state_row_ptr = src_c4_state + layer_i64 * src_c4_state_stride0 + src_state_row * src_c4_state_stride1 - dst_state_row_ptr = dst_c4_state + layer_i64 * dst_c4_state_stride0 + dst_state_row * dst_c4_state_stride1 - src_indexer_row_ptr = ( - src_c4_indexer_state - + layer_i64 * src_c4_indexer_state_stride0 - + src_state_row * src_c4_indexer_state_stride1 - ) - dst_indexer_row_ptr = ( - dst_c4_indexer_state - + layer_i64 * dst_c4_indexer_state_stride0 - + dst_state_row * dst_c4_indexer_state_stride1 - ) - for byte_start in tl.range(0, c4_state_copy_nbytes, BLOCK): + src_pool_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) + dst_pool_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) + src_page = src_pool_slot // swa_pool_page_size + dst_page = dst_pool_slot // swa_pool_page_size + src_page_ptr = src_swa_pool + layer_i64 * dst_swa_pool_stride0 + src_page * dst_swa_pool_stride1 + dst_page_ptr = dst_swa_pool + layer_i64 * dst_swa_pool_stride0 + dst_page * dst_swa_pool_stride1 + for byte_start in tl.range(0, swa_pool_page_nbytes, BLOCK): offsets = byte_start + lanes offsets_i64 = offsets.to(tl.int64) - state_mask = offsets < c4_state_row_nbytes - tl.store( - dst_state_row_ptr + offsets_i64, - tl.load(src_state_row_ptr + offsets_i64, mask=state_mask), - mask=state_mask, + mask = offsets < swa_pool_page_nbytes + tl.store(dst_page_ptr + offsets_i64, tl.load(src_page_ptr + offsets_i64, mask=mask), mask=mask) + + if HAS_C4: + state_pid = task_pid - swa_program_num + if state_pid >= 0: + row = state_pid % 4 + layer = state_pid // 4 + row_i64 = row.to(tl.int64) + layer_i64 = layer.to(tl.int64) + src_c4_state = tl.load(source_ptr_row + 7).to(tl.pointer_type(tl.uint8)) + src_c4_indexer_state = tl.load(source_ptr_row + 8).to(tl.pointer_type(tl.uint8)) + full_offset = token_num - 4 + row_i64 + src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) + dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) + src_swa_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) + dst_swa_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) + src_state_row = (src_swa_slot // swa_pool_page_size) * c4_state_ring + src_swa_slot % c4_state_ring + dst_state_row = (dst_swa_slot // swa_pool_page_size) * c4_state_ring + dst_swa_slot % c4_state_ring + + src_state_row_ptr = ( + src_c4_state + layer_i64 * dst_c4_state_stride0 + src_state_row * dst_c4_state_stride1 ) - indexer_mask = offsets < c4_indexer_state_row_nbytes - tl.store( - dst_indexer_row_ptr + offsets_i64, - tl.load(src_indexer_row_ptr + offsets_i64, mask=indexer_mask), - mask=indexer_mask, + dst_state_row_ptr = ( + dst_c4_state + layer_i64 * dst_c4_state_stride0 + dst_state_row * dst_c4_state_stride1 ) + src_indexer_row_ptr = ( + src_c4_indexer_state + + layer_i64 * dst_c4_indexer_state_stride0 + + src_state_row * dst_c4_indexer_state_stride1 + ) + dst_indexer_row_ptr = ( + dst_c4_indexer_state + + layer_i64 * dst_c4_indexer_state_stride0 + + dst_state_row * dst_c4_indexer_state_stride1 + ) + for byte_start in tl.range(0, c4_state_copy_nbytes, BLOCK): + offsets = byte_start + lanes + offsets_i64 = offsets.to(tl.int64) + state_mask = offsets < c4_state_row_nbytes + tl.store( + dst_state_row_ptr + offsets_i64, + tl.load(src_state_row_ptr + offsets_i64, mask=state_mask), + mask=state_mask, + ) + indexer_mask = offsets < c4_indexer_state_row_nbytes + tl.store( + dst_indexer_row_ptr + offsets_i64, + tl.load(src_indexer_row_ptr + offsets_i64, mask=indexer_mask), + mask=indexer_mask, + ) -def copy_dsv4_dp_cache( - src_mem_manager, +def copy_dsv4_dp_caches( + source_pool_ptrs: torch.Tensor, dst_mem_manager, - src_full_slots: torch.Tensor, - dst_full_slots: torch.Tensor, + task_meta: torch.Tensor, + history_meta: torch.Tensor, ) -> None: - """Copy one 256-token-aligned DP radix suffix directly between DSV4 GPU pools.""" - src_full_slots = src_full_slots.reshape(-1) - dst_full_slots = dst_full_slots.reshape(-1) - token_num = src_full_slots.numel() - assert ( - token_num == dst_full_slots.numel() - and token_num >= DSV4_PROMPT_CACHE_PAGE_SIZE - and token_num % DSV4_PROMPT_CACHE_PAGE_SIZE == 0 - ) - - has_c4 = dst_mem_manager.c4_pool is not None - has_c128 = dst_mem_manager.c128_pool is not None + """Copy aligned DP suffixes; history_meta rows are (task index, local 256-token block).""" + task_num = task_meta.numel() // _TASK_META_WIDTH history_layer_num = max(dst_mem_manager.n_c4, dst_mem_manager.n_c128) - history_program_num = token_num // DSV4_PROMPT_CACHE_PAGE_SIZE * history_layer_num + history_program_num = history_meta.numel() // 2 * history_layer_num swa_program_num = dst_mem_manager.layer_num * 2 c4_state_program_num = dst_mem_manager.n_c4 * 4 - src_c4_pool = src_mem_manager.c4_pool.buffer if has_c4 else None + has_c4 = dst_mem_manager.c4_pool is not None + has_c128 = dst_mem_manager.c128_pool is not None dst_c4_pool = dst_mem_manager.c4_pool.buffer if has_c4 else None - src_c4_indexer_pool = src_mem_manager.c4_indexer_pool.buffer if has_c4 else None dst_c4_indexer_pool = dst_mem_manager.c4_indexer_pool.buffer if has_c4 else None - src_c128_pool = src_mem_manager.c128_pool.buffer if has_c128 else None dst_c128_pool = dst_mem_manager.c128_pool.buffer if has_c128 else None - src_swa_pool = src_mem_manager.swa_pool.buffer dst_swa_pool = dst_mem_manager.swa_pool.buffer - src_c4_state = src_mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None dst_c4_state = dst_mem_manager.c4_state_buffer.view(torch.uint8) if has_c4 else None - src_c4_indexer_state = src_mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None dst_c4_indexer_state = dst_mem_manager.c4_indexer_state_buffer.view(torch.uint8) if has_c4 else None c4_page_nbytes = dst_mem_manager.c4_pool.bytes_per_page if has_c4 else 0 c4_indexer_page_nbytes = dst_mem_manager.c4_indexer_pool.bytes_per_page if has_c4 else 0 c4_state_row_nbytes = dst_c4_state.shape[-1] if has_c4 else 0 c4_indexer_state_row_nbytes = dst_c4_indexer_state.shape[-1] if has_c4 else 0 + program_num = history_program_num + task_num * (swa_program_num + c4_state_program_num) - _copy_dsv4_dp_cache_kernel[(history_program_num + swa_program_num + c4_state_program_num,)]( - src_full_slots, - dst_full_slots, - src_mem_manager.full_to_c4_indexs if has_c4 else None, + _copy_dsv4_dp_caches_kernel[(program_num,)]( + source_pool_ptrs, + task_meta, + history_meta, dst_mem_manager.full_to_c4_indexs if has_c4 else None, - src_c4_pool, - src_c4_pool.stride(0) if has_c4 else 0, - src_c4_pool.stride(1) if has_c4 else 0, dst_c4_pool, dst_c4_pool.stride(0) if has_c4 else 0, dst_c4_pool.stride(1) if has_c4 else 0, - src_c4_indexer_pool, - src_c4_indexer_pool.stride(0) if has_c4 else 0, - src_c4_indexer_pool.stride(1) if has_c4 else 0, dst_c4_indexer_pool, dst_c4_indexer_pool.stride(0) if has_c4 else 0, dst_c4_indexer_pool.stride(1) if has_c4 else 0, - src_mem_manager.full_to_c128_indexs if has_c128 else None, dst_mem_manager.full_to_c128_indexs if has_c128 else None, - src_c128_pool, - src_c128_pool.stride(0) if has_c128 else 0, - src_c128_pool.stride(1) if has_c128 else 0, dst_c128_pool, dst_c128_pool.stride(0) if has_c128 else 0, dst_c128_pool.stride(1) if has_c128 else 0, - src_mem_manager.full_to_swa_indexs, dst_mem_manager.full_to_swa_indexs, - src_swa_pool, - src_swa_pool.stride(0), - src_swa_pool.stride(1), dst_swa_pool, dst_swa_pool.stride(0), dst_swa_pool.stride(1), - src_c4_state, - src_c4_state.stride(0) if has_c4 else 0, - src_c4_state.stride(1) if has_c4 else 0, dst_c4_state, dst_c4_state.stride(0) if has_c4 else 0, dst_c4_state.stride(1) if has_c4 else 0, - src_c4_indexer_state, - src_c4_indexer_state.stride(0) if has_c4 else 0, - src_c4_indexer_state.stride(1) if has_c4 else 0, dst_c4_indexer_state, dst_c4_indexer_state.stride(0) if has_c4 else 0, dst_c4_indexer_state.stride(1) if has_c4 else 0, - token_num, history_program_num, + task_num, + source_pool_ptr_count=_SOURCE_POOL_PTR_COUNT, + task_meta_width=_TASK_META_WIDTH, history_block_size=DSV4_PROMPT_CACHE_PAGE_SIZE, history_layer_num=history_layer_num, c4_ratio=_C4_RATIO, @@ -333,8 +316,7 @@ def copy_dsv4_dp_cache( swa_pool_page_size=dst_mem_manager.swa_pool.page_size, swa_pool_page_nbytes=dst_mem_manager.swa_pool.bytes_per_page, swa_program_num=swa_program_num, - src_c4_state_ring=src_mem_manager.c4_state_ring, - dst_c4_state_ring=dst_mem_manager.c4_state_ring, + c4_state_ring=dst_mem_manager.c4_state_ring, c4_state_row_nbytes=c4_state_row_nbytes, c4_indexer_state_row_nbytes=c4_indexer_state_row_nbytes, c4_state_copy_nbytes=max(c4_state_row_nbytes, c4_indexer_state_row_nbytes), 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 7d2835a232..9c9ba9fdb0 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -305,6 +305,8 @@ def init_dp_kv_shared(self): self.mem_managers.append(MemoryManager.loads_from_shm(rank_idx)) else: self.mem_managers.append(self.model.mem_manager) + if self.is_deepseek_v4: + self.dp_kv_shared_module.init_dsv4_cache_transfer(self.mem_managers) return def get_max_total_token_num(self): diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index f982f4ede5..a93c7cb7f8 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -11,8 +11,9 @@ from ...infer_batch import InferReq from lightllm.utils.dist_utils import get_current_device_id 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 import torch.distributed as dist -from lightllm.models.deepseek_v4.triton_kernel.dp_cache_io import copy_dsv4_dp_cache +from lightllm.models.deepseek_v4.triton_kernel.dp_cache_io import copy_dsv4_dp_caches class DPKVSharedMoudle: @@ -40,6 +41,27 @@ def __init__(self, max_req_num: int, dp_size_in_node: int, backend): assert kv_trans_use_p2p(), "DeepSeek-V4 DP prompt-cache fetch requires P2P KV transfer" + def init_dsv4_cache_transfer(self, mem_managers: List[MemoryManager]) -> None: + pointer_rows = [] + for mem_manager in mem_managers: + has_c4 = mem_manager.c4_pool is not None + has_c128 = mem_manager.c128_pool is not None + pointer_rows.append( + [ + mem_manager.full_to_c4_indexs.data_ptr() if has_c4 else 0, + mem_manager.c4_pool.buffer.data_ptr() if has_c4 else 0, + mem_manager.c4_indexer_pool.buffer.data_ptr() if has_c4 else 0, + mem_manager.full_to_c128_indexs.data_ptr() if has_c128 else 0, + mem_manager.c128_pool.buffer.data_ptr() if has_c128 else 0, + mem_manager.full_to_swa_indexs.data_ptr(), + mem_manager.swa_pool.buffer.data_ptr(), + mem_manager.c4_state_buffer.data_ptr() if has_c4 else 0, + mem_manager.c4_indexer_state_buffer.data_ptr() if has_c4 else 0, + ] + ) + self.dsv4_source_pool_ptrs = torch.tensor(pointer_rows, dtype=torch.uint64, device="cuda") + return + def fill_reqs_info(self, reqs: List[InferReq]): """ 填充请求的 kv 信息到共享内存中 @@ -136,26 +158,69 @@ def kv_trans(self, trans_tasks: List["TransTask"]): if len(trans_tasks) > 0: if self.backend.is_deepseek_v4: req_manager = g_infer_context.req_manager + prompt_cache_page_size = req_manager.get_prompt_cache_page_size() + req_list = [] + ready_list = [] + seq_list = [] + dst_full_slot_views = [] + task_meta_data = [] + history_block_nums = [] for trans_task in trans_tasks: start = trans_task.req.cur_kv_len end = start + len(trans_task.mem_indexes) dst_full_slots = req_manager.req_to_token_indexs[trans_task.req.req_idx, start:end] dst_full_slots.copy_(trans_task.mem_indexes, non_blocking=True) - req_manager.prepare_pd_decode_cache( - req_idx=trans_task.req.req_idx, - ready_cache_len=start, - input_len=end, - new_full_slots=dst_full_slots, + req_list.append(trans_task.req.req_idx) + ready_list.append(start) + seq_list.append(end) + dst_full_slot_views.append(dst_full_slots) + task_meta_data.extend( + [ + trans_task.max_kv_len_mem_manager_index, + end - start, + trans_task.max_kv_len_mem_indexes.data_ptr(), + dst_full_slots.data_ptr(), + ] ) + history_block_nums.append((end - start) // prompt_cache_page_size) + + # Keep full slots in the same request-major order as req_list. + new_full_slots = ( + dst_full_slot_views[0] if len(dst_full_slot_views) == 1 else torch.cat(dst_full_slot_views) + ) + req_manager.prepare_pd_decode_cache( + req_list=req_list, + ready_list=ready_list, + seq_list=seq_list, + new_full_slots=new_full_slots, + ) + + # The history kernel consumes (task index, block index) pairs in block-major order. + history_meta_data = [] + for block_index in range(max(history_block_nums)): + for task_index, block_num in enumerate(history_block_nums): + if block_index < block_num: + history_meta_data.extend([task_index, block_index]) + + # transfer_meta packs two flat uint64 tables into one H2D copy: + # task_meta: (source manager, token count, source slots pointer, destination slots pointer) + # history_meta: (task index, block index) + # For two tasks with 2 and 1 history blocks, the layout is: + # [task0 fields, task1 fields, (0, 0), (1, 0), (0, 1)] + task_meta_size = len(task_meta_data) + transfer_meta = g_pin_mem_manager.gen_from_list( + key="dsv4_dp_cache_transfer_meta", + data=task_meta_data + history_meta_data, + dtype=torch.uint64, + ).to(req_manager.req_to_token_indexs.device, non_blocking=True) + dst_mem_manager = self.backend.model.mem_manager - dst_req_to_token = self.backend.model.req_manager.req_to_token_indexs - for trans_task in trans_tasks: - src_mem_manager = self.backend.mem_managers[trans_task.max_kv_len_mem_manager_index] - start = trans_task.req.cur_kv_len - end = start + len(trans_task.mem_indexes) - src_mem_indexes = trans_task.max_kv_len_mem_indexes - dst_mem_indexes = dst_req_to_token[trans_task.req.req_idx, start:end] - copy_dsv4_dp_cache(src_mem_manager, dst_mem_manager, src_mem_indexes, dst_mem_indexes) + copy_dsv4_dp_caches( + source_pool_ptrs=self.dsv4_source_pool_ptrs, + dst_mem_manager=dst_mem_manager, + task_meta=transfer_meta[:task_meta_size], + history_meta=transfer_meta[task_meta_size:], + ) else: max_kv_len_mem_indexes = [] max_kv_len_dp_ranks = [] @@ -181,8 +246,8 @@ def kv_trans(self, trans_tasks: List["TransTask"]): transfer_token_num = sum(len(trans_task.mem_indexes) for trans_task in trans_tasks) self.backend.logger.info(f"dp_i {self.dp_rank_in_node} transfer kv tokens num: {transfer_token_num}") - if self.backend.is_deepseek_v4: - # Source radix mappings stay alive until every peer copy on the current stream finishes. + if self.backend.is_deepseek_v4 and self.backend.args.enable_cpu_cache: + # CPU-cache restore can evict source radix pages before the scheduler all-gather fences this stream. dist.barrier(group=self.backend.node_nccl_group) for trans_task in trans_tasks: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index cca2d0791d..bd893a2b2b 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -135,9 +135,9 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): ] = mem_indexes if isinstance(req_manager, DeepseekV4ReqManager): req_manager.prepare_pd_decode_cache( - req_idx=req_obj.req_idx, - ready_cache_len=req_obj.cur_kv_len, - input_len=input_len, + req_list=[req_obj.req_idx], + ready_list=[req_obj.cur_kv_len], + seq_list=[input_len], new_full_slots=mem_indexes, ) From b46551cf8a7f729c1965f8660e720588ef06d907 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 11 Aug 2026 01:32:13 +0000 Subject: [PATCH 104/214] fix dsv4 dpep2 --- lightllm/distributed/communication_op.py | 2 ++ lightllm/models/deepseek_v4/model.py | 9 +++++---- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index d003f3f3a1..1838828cb1 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -148,6 +148,8 @@ def get_moe_quant_methods(layer_weights: List) -> Set[str]: for layer_weight in layer_weights: # dense 层没有 experts;这里只关心真正参与 MoE 计算的层。 experts = getattr(layer_weight, "experts", None) + if experts is None: + experts = getattr(layer_weight, "experts_", None) quant_method = getattr(experts, "quant_method", None) method_name = getattr(quant_method, "method_name", None) if method_name is not None: diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 1022e0c03d..4252698587 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -144,10 +144,11 @@ def _init_custom(self): for layer in self.layers_infer: layer.dsv4_prefill_aux_stream = prefill_aux_stream dist_group_manager.new_deepep_group( - self.config["n_routed_experts"], - self.config["hidden_size"], - self.config.get("num_experts_per_tok", 1), - self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), + n_routed_experts=self.config["n_routed_experts"], + hidden_size=self.config["hidden_size"], + expert_quant_method_names=dist_group_manager.get_moe_quant_methods(self.trans_layers_weight), + num_experts_per_tok=self.config.get("num_experts_per_tok", 1), + moe_intermediate_size=self.config.get("moe_intermediate_size", self.config.get("intermediate_size")), ) return From 2cbf6624d14b1ad824b3c1275fdcde9468aa95af Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 11 Aug 2026 05:50:40 +0000 Subject: [PATCH 105/214] disable c4 aux stream during prefill microbatch overlap --- lightllm/models/deepseek_v4/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 4252698587..ea0d9cbbb5 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -139,7 +139,7 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() self.dsv4_workspace = DeepseekV4Workspace(self) - if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1": + if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1" and not self.args.enable_prefill_microbatch_overlap: prefill_aux_stream = torch.cuda.Stream() for layer in self.layers_infer: layer.dsv4_prefill_aux_stream = prefill_aux_stream From eae1968000f6446442f2dbe1d64ff112a75bc276 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 11 Aug 2026 07:23:42 +0000 Subject: [PATCH 106/214] fix pdl cuda bug --- .../deepseek_v4/triton_kernel/csrc/norm_rope.cu | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu index b0a2963a53..a2aadcdd72 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu @@ -84,6 +84,7 @@ template __device__ __forceinline__ void pdl_wait_primary() { #if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 if constexpr (UsePDL) { + // PDL may start this grid early, so this must precede its first global load. asm volatile("griddepcontrol.wait;" ::: "memory"); } #endif @@ -176,6 +177,8 @@ __global__ __launch_bounds__(kFusedQBlockSize, 16) void fused_q_norm_rope_kernel const uint32_t total_works = params.batch_size * params.num_q_heads; if (work_id >= total_works) return; + pdl_wait_primary(); + const uint32_t batch_id = work_id / params.num_q_heads; const uint32_t head_id = work_id % params.num_q_heads; const auto* input_ptr = @@ -186,8 +189,6 @@ __global__ __launch_bounds__(kFusedQBlockSize, 16) void fused_q_norm_rope_kernel __shared__ Storage rope_storage[kFusedQNumWarps][kRopeVecs]; - pdl_wait_primary(); - Storage input_vec[kLocalSize]; #pragma unroll for (int i = 0; i < kLocalSize; ++i) { @@ -260,13 +261,13 @@ __global__ __launch_bounds__(kFusedKBlockSize, 8) void fused_k_norm_rope_flashml const uint32_t work_id = blockIdx.x; if (work_id >= params.batch_size) return; + pdl_wait_primary(); + const auto* input_ptr = params.kv + work_id * params.kv_stride_batch; const int32_t position = static_cast(static_cast(params.positions)[work_id]); const int32_t out_loc = params.out_loc[work_id]; const float* freqs_cis = params.freqs_cis + position * kMainRopeDim; - pdl_wait_primary(); - const Storage input_vec = reinterpret_cast(input_ptr)[tx]; const Storage weight_vec = reinterpret_cast(params.kv_weight)[tx]; float2 data; @@ -341,14 +342,14 @@ __global__ __launch_bounds__(kFusedQBlockSize, 16) void fused_q_indexer_rope_had const uint32_t total_works = params.batch_size * params.num_heads; if (work_id >= total_works) return; + pdl_wait_primary(); + const uint32_t batch_id = work_id / params.num_heads; const int32_t position = static_cast(static_cast(params.positions)[batch_id]); const auto* input_ptr = params.q_input + static_cast(work_id) * kIndexerHeadDim; const float* freqs_cis = params.freqs_cis + position * kIndexerRopeDim; const bool is_rope_lane = lane_id >= kWarpThreads - kRopeVecs; - pdl_wait_primary(); - const float weight_value = __bfloat162float(params.weight[work_id]); const Storage input_vec = reinterpret_cast(input_ptr)[lane_id]; Float4 data; From f9e93b9210eee0d80a8e1c1ce752f47583506b07 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 12 Aug 2026 01:41:43 +0000 Subject: [PATCH 107/214] pack need owner gpu --- .../pd/decode_node_impl/decode_impl.py | 14 +++++++++----- .../pd/prefill_node_impl/prefill_impl.py | 12 ++++++++---- 2 files changed, 17 insertions(+), 9 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index bd893a2b2b..1c2828f89f 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -200,11 +200,15 @@ def _create_pd_trans_task( ): # 确定传输设备 if req_obj.pd_trans_device_id == -1: - if not hasattr(self, "pd_iter_device_id"): - self.pd_iter_device_id = 0 - req_obj.pd_trans_device_id = self.pd_iter_device_id - # only self.is_master_in_dp will be used. - self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size + if self.is_deepseek_v4: + # DSV4 packed cache belongs to this DP rank; its unpack kernel must run on the owner GPU. + req_obj.pd_trans_device_id = self.dp_rank_in_node + else: + if not hasattr(self, "pd_iter_device_id"): + self.pd_iter_device_id = 0 + req_obj.pd_trans_device_id = self.pd_iter_device_id + # only self.is_master_in_dp will be used. + self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size if page_kind not in ("kv", "linear_att_state"): raise ValueError(f"unknown PD trans page kind {page_kind}") diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index ede4a9c865..39bf5f5751 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -100,10 +100,14 @@ def _create_pd_trans_task( ) -> PDChunckedTransTask: # 确定传输设备 if req_obj.pd_trans_device_id == -1: - if not hasattr(self, "pd_iter_device_id"): - self.pd_iter_device_id = 0 - req_obj.pd_trans_device_id = self.pd_iter_device_id - self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size + if self.is_deepseek_v4: + # DSV4 packed cache belongs to this DP rank; its pack kernel must run on the owner GPU. + req_obj.pd_trans_device_id = self.dp_rank_in_node + else: + if not hasattr(self, "pd_iter_device_id"): + self.pd_iter_device_id = 0 + req_obj.pd_trans_device_id = self.pd_iter_device_id + self.pd_iter_device_id = (self.pd_iter_device_id + 1) % self.node_world_size pd_decode_node_info = req_obj.sampling_param.pd_decode_node if page_kind == "kv": From 7958b80e4fe68ef1bdb43e6fa000a6aa454516ff Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 12 Aug 2026 03:31:25 +0000 Subject: [PATCH 108/214] pd try fix --- lightllm/server/httpserver/pd_loop.py | 26 +++++- .../mode_backend/pd/nixl_kv_transporter.py | 86 +++++++++++-------- 2 files changed, 73 insertions(+), 39 deletions(-) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index b1b99a3663..0c798c9535 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -90,6 +90,8 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O uri, max_size=get_lightllm_websocket_max_message_size(), max_queue=(2048 * 1024, 2048 * 1023), # 关键修改 + # 下方应用层心跳已负责存活检测,禁用协议层 keepalive,避免繁忙连接被误断。 + ping_interval=None, ) as websocket: sock = websocket.transport.get_extra_info("socket") @@ -239,12 +241,34 @@ async def _pd_process_generate( # 转发token的task async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): + max_message_size = get_lightllm_websocket_max_message_size() + while True: handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: load_info: dict = _get_load_info() - await websocket.send(pickle.dumps((ObjType.TOKEN_PACKS, handle_list, load_info))) + pending_handle_lists = [handle_list] + while pending_handle_lists: + token_list = pending_handle_lists.pop() + payload = pickle.dumps((ObjType.TOKEN_PACKS, token_list, load_info)) + if len(payload) <= max_message_size: + await websocket.send(payload) + continue + + if len(token_list) == 1: + raise ValueError( + f"single PD token pack is {len(payload)} bytes, exceeding websocket limit " + f"{max_message_size}" + ) + + split_index = len(token_list) // 2 + logger.warning( + f"PD token pack is {len(payload)} bytes with {len(token_list)} items, exceeding websocket " + f"limit {max_message_size}; splitting it" + ) + # 栈后进先出,先压后半段,保持 token 的原始发送顺序。 + pending_handle_lists.extend((token_list[split_index:], token_list[:split_index])) async def _send_heartbeat_to_pd_master(websocket: ClientConnection): diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index 55ff077665..6d43b0aa64 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -1,6 +1,7 @@ import pickle import copy import os +import threading import time from dataclasses import dataclass from typing import Dict @@ -15,6 +16,7 @@ from nixl._api import nixl_agent as NixlWrapper from nixl._api import nixlBind from nixl._api import nixl_agent_config + from nixl._api import nixl_thread_sync_t logger.info("Nixl is available") except ImportError: @@ -32,13 +34,13 @@ def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): "yes", "on", ) - conf = None + conf = nixl_agent_config(sync_mode=nixl_thread_sync_t.NIXL_THREAD_SYNC_RW) if self.capture_telemetry: - conf = nixl_agent_config() conf.capture_telemetry = True logger.info("NIXL telemetry enabled") self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) + self._remote_agents_lock = threading.Lock() self.remote_agents: Dict[str, PDAgentMetadata] = {} return @@ -99,50 +101,53 @@ def _get_remote_page_xfer_handles(self, remote_agent: PDAgentMetadata, transfer_ return remote_agent.page_xfer_handles[transfer_nbytes] def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): - if remote_agent.agent_name in self.remote_agents: - return + with self._remote_agents_lock: + if remote_agent.agent_name in self.remote_agents: + return - start_time = time.time() + start_time = time.time() - peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) - if isinstance(peer_name, bytes): - peer_name = peer_name.decode() + peer_name = self.nixl_agent.add_remote_agent(remote_agent.agent_metadata) + if isinstance(peer_name, bytes): + peer_name = peer_name.decode() - assert ( - peer_name == remote_agent.agent_name - ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" + assert ( + peer_name == remote_agent.agent_name + ), f"Peer name {peer_name} does not match remote name {remote_agent.agent_name}" - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - remote_agent.page_xfer_handles = { - self.page_len: self._create_paged_xfer_handles( - page_mem_desc, - remote_agent.num_pages, - self.page_len, - agent_name=peer_name, + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + remote_agent.page_xfer_handles = { + self.page_len: self._create_paged_xfer_handles( + page_mem_desc, + remote_agent.num_pages, + self.page_len, + agent_name=peer_name, + ) + } + + logger.info( + f"Added remote agent {peer_name} with mem desc {page_mem_desc} " + f"cost time: {time.time() - start_time} s" ) - } - logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" - ) - - self.remote_agents[remote_agent.agent_name] = remote_agent + self.remote_agents[remote_agent.agent_name] = remote_agent return def remove_remote_agent(self, peer_name: str): - if peer_name in self.remote_agents: - try: - remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) - assert remote_agent.agent_name == peer_name - self.nixl_agent.remove_remote_agent(remote_agent.agent_name) - if remote_agent.page_xfer_handles is not None: - for handles in remote_agent.page_xfer_handles.values(): - self.nixl_agent.release_dlist_handle(handles) - except BaseException as e: - logger.error(f"remove remote agent {peer_name} failed") - logger.exception(str(e)) - else: - logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") + with self._remote_agents_lock: + if peer_name in self.remote_agents: + try: + remote_agent: PDAgentMetadata = self.remote_agents.pop(peer_name, None) + assert remote_agent.agent_name == peer_name + self.nixl_agent.remove_remote_agent(remote_agent.agent_name) + if remote_agent.page_xfer_handles is not None: + for handles in remote_agent.page_xfer_handles.values(): + self.nixl_agent.release_dlist_handle(handles) + except BaseException as e: + logger.error(f"remove remote agent {peer_name} failed") + logger.exception(str(e)) + else: + logger.warning(f"try to remove remote agent, but peer name {peer_name} agent did not exist") def send_write_done_task_to_decode_node(self, trans_task: PDChunckedTransTask): decode_agent_name = trans_task.decode_agent_name @@ -302,7 +307,12 @@ def write_blocks_paged( def check_task_status(self, trans_task: PDChunckedTransTask) -> str: assert trans_task.xfer_handle is not None handle = trans_task.xfer_handle - xfer_state = self.nixl_agent.check_xfer_state(handle) + try: + xfer_state = self.nixl_agent.check_xfer_state(handle) + except Exception as e: + logger.error(f"Check transfer state failed with trans task {trans_task.to_str()} for handle {handle}") + logger.exception(str(e)) + return "ERR" if xfer_state == "ERR": logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") return xfer_state From a0ea9b2bed069b5f793d184d1cfbf65cf0b56be6 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 13 Aug 2026 07:06:51 +0000 Subject: [PATCH 109/214] async execute self.tokens --- lightllm/server/httpserver_for_pd_master/manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 18cfa0d51f..04556d8423 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -152,7 +152,7 @@ async def _generate( await multimodal_params.verify_and_preload(request) # 计算输入的 input_token_num, 进行校验,如果输入+输出参数设置太长,则将 # sampling_params 的参数进行修正。 - input_token_num = self.tokens(prompt, multimodal_params, sampling_params) + input_token_num = await asyncio.to_thread(self.tokens, prompt, multimodal_params, sampling_params) fake_prompt_ids = [0 for _ in range(input_token_num)] from lightllm.server.httpserver.manager import HttpServerManager From fc8b7ec411774fab269e1e3799efff4ac15f826e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 17 Aug 2026 02:42:29 +0000 Subject: [PATCH 110/214] warmup tilelang --- lightllm/common/basemodel/basemodel.py | 5 ++ .../layer_infer/hyper_connection.py | 14 ++++ lightllm/models/deepseek_v4/model.py | 64 +++++++++++++++++++ 3 files changed, 83 insertions(+) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index a92f8687e0..885a02f8dc 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -140,6 +140,7 @@ def __init__(self, kvargs): self._autotune_warmup() self._full_att_decode_autotune() + self._kernel_warmup() self._init_padded_req() self._init_cudagraph() self._init_prefill_cuda_graph() @@ -350,6 +351,10 @@ def _full_att_decode_autotune(self): def _init_custom(self): pass + def _kernel_warmup(self): + """Warm model-specific kernels before CUDA graph capture.""" + return + @torch.no_grad() def forward(self, model_input: ModelInput): model_input.to_cuda() diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index c2592c4d0b..920309f53d 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -1,4 +1,5 @@ import torch +from vllm._tilelang_ops import compute_num_split try: import vllm.model_executor.layers.mhc # noqa: F401 @@ -10,6 +11,19 @@ HC_POST_ALPHA = 2.0 +def mhc_warmup_token_sizes(max_tokens, hidden_size, hc_mult): + """Choose one token count for every reachable mHC split-K variant.""" + block_size = 64 + token_sizes_by_split = {} + + max_grid_size = (max_tokens + block_size - 1) // block_size + for grid_size in range(1, max_grid_size + 1): + n_splits = compute_num_split(block_size, hc_mult * hidden_size, grid_size) + token_sizes_by_split.setdefault(n_splits, min(grid_size * block_size, max_tokens)) + + return sorted(token_sizes_by_split.values()) + + def hc_pre(residual, hc_fn, hc_scale, hc_base, rms_eps, hc_eps, sinkhorn_iters, norm_weight, norm_eps): """Standalone hc_pre for the first layer. residual:[T, hc, dim] -> (x[T,dim], residual, post_mix[T,hc,1], res_mix[T,hc,hc]); the sub-layer RMSNorm is fused via norm_weight.""" diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index ea0d9cbbb5..bae2e85a55 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -2,6 +2,7 @@ import importlib.util import json import os +import time import torch from lightllm.models.registry import ModelRegistry @@ -29,6 +30,11 @@ from lightllm.common.basemodel.attention.nsa.dsv4_fp8_flashmla_sparse import DSV4_NSA_BACKENDS from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo from lightllm.models.deepseek_v4.workspace import DeepseekV4Workspace +from lightllm.models.deepseek_v4.layer_infer.hyper_connection import ( + hc_head, + hc_post, + mhc_warmup_token_sizes, +) from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, linear_ramp_mask, @@ -152,6 +158,64 @@ def _init_custom(self): ) return + @torch.no_grad() + def _kernel_warmup(self): + if self.is_mtp_draft_model: + return + + layer_infer = self.layers_infer[0] + layer_weight = self.trans_layers_weight[0] + hidden_size = self.config["hidden_size"] + hc_mult = self.config["hc_mult"] + split_token_sizes = mhc_warmup_token_sizes( + max_tokens=self.batch_max_tokens, + hidden_size=hidden_size, + hc_mult=hc_mult, + ) + token_sizes = sorted(set(split_token_sizes + [size for size in (1, 8, 17) if size <= self.batch_max_tokens])) + + started = time.perf_counter() + logger.info( + "warming DeepSeek-V4 mHC TileLang kernels for token sizes %s", + token_sizes, + ) + residual = torch.zeros( + max(token_sizes), + hc_mult, + hidden_size, + dtype=torch.bfloat16, + device="cuda", + ) + for token_size in split_token_sizes: + layer_infer._hc_attn_in(residual[:token_size], layer_weight) + + for token_size in (size for size in (1, 8, 17) if size <= self.batch_max_tokens): + hc_state = layer_infer._hc_attn_in(residual[:token_size], layer_weight) + hc_state = layer_infer._hc_ffn_in(*hc_state, layer_weight) + + streams = hc_post(*hc_state) + hc_head( + streams, + self.pre_post_weight.hc_head_fn_.weight, + self.pre_post_weight.hc_head_scale_.weight, + self.pre_post_weight.hc_head_base_.weight, + hc_mult, + hidden_size, + self.config["rms_norm_eps"], + self.config.get("hc_eps", 1e-6), + torch.empty, + ) + + torch.cuda.synchronize() + del residual, hc_state, streams + torch.cuda.empty_cache() + logger.info( + "DeepSeek-V4 mHC TileLang warmup finished in %.2f seconds (%d split-K variants)", + time.perf_counter() - started, + len(split_token_sizes), + ) + return + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: """Commit DSV4 derived slots before BaseModel pads or scatters the generic input.""" if model_input.is_prefill and self.is_mtp_draft_model: From 082c4590b3d17fdede13a7cad206bfb4f383e8f7 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 17 Aug 2026 08:01:35 +0000 Subject: [PATCH 111/214] default thinking high --- lightllm/models/deepseek_v4/model.py | 10 +++++- lightllm/server/api_openai.py | 2 +- lightllm/utils/config_utils.py | 48 +++++++++++++++++++++++++++- 3 files changed, 57 insertions(+), 3 deletions(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index bae2e85a55..0da693f6bd 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -40,6 +40,7 @@ linear_ramp_mask, ) from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num, get_env_start_args +from lightllm.utils.config_utils import normalize_deepseek_v4_config from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager @@ -58,6 +59,11 @@ class DeepseekV4TpPartModel(LlamaTpPartModel): post_layer_infer_class = DeepseekV4PostLayerInfer transformer_layer_infer_class = DeepseekV4TransformerLayerInfer + def _init_config(self): + super()._init_config() + normalize_deepseek_v4_config(self.config) + return + infer_state_class = DeepseekV4InferStateInfo def _verify_params(self): @@ -423,9 +429,11 @@ def apply_chat_template( msgs.insert(0, {"role": "system", "content": "", "tools": wrapped_tools}) if thinking is None: - thinking = bool(enable_thinking) if enable_thinking is not None else False + thinking = bool(enable_thinking) if enable_thinking is not None else True thinking_mode = "thinking" if thinking else "chat" effort = kwargs.get("reasoning_effort") + if thinking and effort is None: + effort = "high" if effort not in ("max", "high", None): effort = None encoding = self._get_encoding_module() diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index b070750c40..d825266206 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -163,7 +163,7 @@ def _is_force_thinking_mode(request: ChatCompletionRequest) -> bool: return chat_template_kwargs["thinking"] is True if request.reasoning_effort is not None: return request.reasoning_effort != "none" - return False + return True if reasoning_parser in ["qwen3", "glm45", "nano_v3", "interns1", "gemma4"]: # qwen3, glm45, nano_v3, interns1, and gemma4 are reasoning by default; return not request.chat_template_kwargs or request.chat_template_kwargs.get("enable_thinking", True) is True diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 4b49cabd66..7b2da8308a 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -8,10 +8,56 @@ logger = init_logger(__name__) +def normalize_deepseek_v4_config(config: Dict[str, Any]) -> Dict[str, Any]: + """Normalize the newer DeepSeek-V4 layer schema for existing runtimes.""" + if config.get("model_type") != "deepseek_v4": + return config + + if "num_hash_layers" not in config: + mlp_layer_types = config.get("mlp_layer_types") + if isinstance(mlp_layer_types, list): + num_hash_layers = sum(layer_type == "hash_moe" for layer_type in mlp_layer_types) + if mlp_layer_types[:num_hash_layers] != ["hash_moe"] * num_hash_layers: + raise ValueError("DeepSeek-V4 hash_moe layers must form a contiguous prefix") + config["num_hash_layers"] = num_hash_layers + + if "rope_scaling" not in config: + rope_parameters = config.get("rope_parameters") + if isinstance(rope_parameters, dict): + compress_rope = rope_parameters.get("compress") + if isinstance(compress_rope, dict): + config["rope_scaling"] = dict(compress_rope) + + if "compress_ratios" in config: + return config + + layer_types = config.get("layer_types") + compress_rates = config.get("compress_rates") + if not isinstance(layer_types, list) or not isinstance(compress_rates, dict): + return config + + ratios = [] + for layer_type in layer_types: + # Uncompressed layer types (for example sliding_attention) are omitted + # from the new mapping and therefore use a compression ratio of zero. + ratio = compress_rates.get(layer_type, 0) + if ratio not in (0, 4, 128): + raise ValueError(f"Unsupported DeepSeek-V4 compression ratio {ratio!r} for layer type {layer_type!r}") + ratios.append(ratio) + + # Older checkpoints include one trailing entry for the optional MTP KV + # layer. Preserve that layout so both regular and MTP paths can slice it. + if len(ratios) == config.get("num_hidden_layers"): + ratios.append(0) + + config["compress_ratios"] = ratios + return config + + def get_config_json(model_path: str): with open(os.path.join(model_path, "config.json"), "r") as file: json_obj = json.load(file) - return json_obj + return normalize_deepseek_v4_config(json_obj) def get_generation_config_diff_dict(model_path: str) -> Dict[str, Any]: From cfce08726b09e44182ff63df2978ae95cfea1337 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 18 Aug 2026 12:28:49 +0000 Subject: [PATCH 112/214] perf(pd): compact token forwarding payloads - avoid duplicating prompt_ids in first-token metadata - omit output logprob metadata when it is not requested - send common token updates using a compact tuple format - preserve legacy packets for optional metadata and logprobs - support ordered mixed compact/legacy batches --- lightllm/server/api_openai.py | 34 ++--- lightllm/server/api_start.py | 2 + lightllm/server/httpserver/async_queue.py | 15 ++- lightllm/server/httpserver/manager.py | 3 +- lightllm/server/httpserver/pd_loop.py | 39 +++++- .../httpserver_for_pd_master/manager.py | 116 ++++++++++-------- lightllm/server/pd_io_struct.py | 86 ++++++++++++- .../decode_node_impl/decode_trans_process.py | 21 ++-- 8 files changed, 225 insertions(+), 91 deletions(-) diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index d825266206..8df9458b7f 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -46,7 +46,6 @@ CompletionChoice, CompletionLogprobs, CompletionStreamResponse, - CompletionStreamChoice, FunctionResponse, ToolCall, UsageInfo, @@ -343,6 +342,9 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req sampling_params = SamplingParams() sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict) + # Chat completions don't expose output-token logprobs. PD nodes use this + # transport-only marker to avoid forwarding unused per-token metadata. + sampling_params.return_output_logprobs = False sampling_params.verify() multimodal_params = MultimodalParams(**multimodal_params_dict) @@ -839,6 +841,7 @@ async def completions_impl(request: CompletionRequest, raw_request: Request) -> sampling_params = SamplingParams() sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict) + sampling_params.return_output_logprobs = request.logprobs is not None sampling_params.verify() # v1/completions does not support multimodal inputs, so we use an empty MultimodalParams @@ -935,19 +938,22 @@ async def stream_results() -> AsyncGenerator[bytes, None]: prompt_str = g_objs.httpserver_manager.tokenizer.decode(prompt, skip_special_tokens=False) output_text = prompt_str + output_text - stream_choice = CompletionStreamChoice( - index=choice_index, - text=output_text, - finish_reason=current_finish_reason, - logprobs=None if request.logprobs is None else {}, - ) - stream_resp = CompletionStreamResponse( - id=group_request_id, - created=created_time, - model=request.model, - choices=[stream_choice], - ) - yield f"data: {json.dumps(stream_resp.model_dump(), ensure_ascii=False)}\n\n" + stream_resp = { + "id": str(group_request_id), + "object": "text_completion", + "created": created_time, + "model": request.model, + "choices": [ + { + "text": output_text, + "index": int(choice_index), + "logprobs": None if request.logprobs is None else {}, + "finish_reason": current_finish_reason, + } + ], + "usage": None, + } + yield f"data: {json.dumps(stream_resp, ensure_ascii=False)}\n\n" usage = UsageInfo( prompt_tokens=prompt_tokens, diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 833e553264..a518daa53f 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -26,6 +26,7 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args +from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup logger = init_logger(__name__) @@ -85,6 +86,7 @@ def _launch_subprocesses(args: StartArgs): if args.enable_cpu_cache: # 生成一个用于创建cpu kv cache的共享内存id。 args.cpu_kv_cache_shm_id = uuid.uuid1().int % 123456789 + register_sysv_shm_for_cleanup(args.cpu_kv_cache_shm_id) if args.enable_multimodal: args.multi_modal_cache_shm_id = uuid.uuid1().int % 123456789 diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index a9f0c9068f..47cfed4c88 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -5,7 +5,6 @@ class AsyncQueue: def __init__(self): self.datas = [] self.event = asyncio.Event() - self.lock = asyncio.Lock() async def wait_to_ready(self): try: @@ -14,15 +13,15 @@ async def wait_to_ready(self): pass async def get_all_data(self): - async with self.lock: - self.event.clear() - ans = self.datas - self.datas = [] - return ans + self.event.clear() + ans = self.datas + self.datas = [] + return ans async def put(self, obj): - async with self.lock: - self.datas.append(obj) + was_empty = not self.datas + self.datas.append(obj) + if was_empty: self.event.set() return diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 195ef17e45..f92ee84ec2 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -10,6 +10,7 @@ import hashlib import datetime import pickle +from array import array from frozendict import frozendict asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) @@ -397,7 +398,7 @@ async def generate( f"pd prefill node upload group_req_id {group_request_id} prompt ids len : {len(prompt_ids)}" ) await pd_upload_websocket.send( - pickle.dumps((ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS, group_request_id, prompt_ids)) + pickle.dumps((ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS, group_request_id, array("i", prompt_ids))) ) try: await asyncio.wait_for(pd_event.wait(), timeout=180) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 0c798c9535..d55631e3ff 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -11,7 +11,12 @@ import sys from typing import Dict, Optional, Union, List from websockets import ClientConnection -from lightllm.server.pd_io_struct import NodeRole, ObjType +from lightllm.server.pd_io_struct import ( + NodeRole, + ObjType, + PD_COMPACT_TOKEN_INFO_LEN, + build_pd_compact_token_info, +) from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger @@ -222,6 +227,7 @@ async def _pd_process_generate( pd_upload_websocket: ClientConnection, pd_event: asyncio.Event, ): + return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) try: async for sub_req_id, request_output, metadata, finish_status in manager.generate( prompt=prompt, @@ -231,7 +237,17 @@ async def _pd_process_generate( pd_upload_websocket=pd_upload_websocket, pd_event=pd_event, ): - metadata["node_mode"] = manager.args.run_mode + metadata.pop("prompt_ids", None) + if not return_output_logprobs: + for key in ("id", "logprob", "cumlogprob", "special", "logprobs"): + metadata.pop(key, None) + if metadata.get("count_output_tokens") == 1: + metadata["node_mode"] = manager.args.run_mode + if not return_output_logprobs: + compact_token_info = build_pd_compact_token_info(sub_req_id, request_output, metadata, finish_status) + if compact_token_info is not None: + await forwarding_queue.put(compact_token_info) + continue await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") @@ -248,10 +264,25 @@ async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: Clien if handle_list: load_info: dict = _get_load_info() - pending_handle_lists = [handle_list] + pending_handle_lists = [] + group_start = 0 + group_is_compact = len(handle_list[0]) == PD_COMPACT_TOKEN_INFO_LEN + for index in range(1, len(handle_list)): + item_is_compact = len(handle_list[index]) == PD_COMPACT_TOKEN_INFO_LEN + if item_is_compact != group_is_compact: + pending_handle_lists.append(handle_list[group_start:index]) + group_start = index + group_is_compact = item_is_compact + pending_handle_lists.append(handle_list[group_start:]) + pending_handle_lists.reverse() while pending_handle_lists: token_list = pending_handle_lists.pop() - payload = pickle.dumps((ObjType.TOKEN_PACKS, token_list, load_info)) + obj_type = ( + ObjType.TOKEN_PACKS_COMPACT + if len(token_list[0]) == PD_COMPACT_TOKEN_INFO_LEN + else ObjType.TOKEN_PACKS + ) + payload = pickle.dumps((obj_type, token_list, load_info)) if len(payload) <= max_message_size: await websocket.send(payload) continue diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 04556d8423..01d9830d52 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -12,7 +12,13 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus -from ..pd_io_struct import PD_Client_Obj, PDUpKVStatus, ObjType, PDDecodeNodeInfo +from ..pd_io_struct import ( + PD_Client_Obj, + PDUpKVStatus, + ObjType, + PDDecodeNodeInfo, + unpack_pd_compact_token_info, +) from lightllm.server.core.objs import SamplingParams, StartArgs from ..multimodal_params import MultimodalParams from ..tokenizer import get_tokenizer @@ -160,7 +166,9 @@ async def _generate( self, prompt_ids=fake_prompt_ids, sampling_params=sampling_params ) + return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) + origin_sampling_params.return_output_logprobs = return_output_logprobs origin_group_request_id = self.id_gen.generate_id() # Record one user request even when it is expanded into multiple independent @@ -174,6 +182,7 @@ async def _generate( generators = [] for choice_index in range(choice_count): choice_sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) + choice_sampling_params.return_output_logprobs = return_output_logprobs choice_sampling_params.n = 1 choice_sampling_params.best_of = 1 generators.append( @@ -228,6 +237,7 @@ async def _generate_one( for iter_index, block_max_new_tokens in enumerate(max_new_tokens_list): sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) + sampling_params.return_output_logprobs = getattr(origin_sampling_params, "return_output_logprobs", True) block_group_request_id = self.id_gen.generate_id() sampling_params.group_request_id = block_group_request_id logger.info(f"pd log gen sub req id {block_group_request_id} for main req id {origin_request_id}") @@ -420,30 +430,33 @@ async def fetch_pd_stream( ) first_token_gen = False + next_disconnect_check = 0.0 while True: await req_status.wait_to_ready() - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, - reason="fetch_pd_stream decode period check network disconnected", - ) - if await req_status.can_read(self.req_id_to_out_inf): - token_list = await req_status.pop_all_tokens() - for sub_req_id, request_output, metadata, finish_status in token_list: - output_index = metadata.get("count_output_tokens") - # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 - if output_index == 1: - if first_token_gen is False: - first_token_gen = True - node_run_mode = metadata.pop("node_mode", None) - if node_run_mode == "prefill": - if old_max_new_tokens != 1 and finish_status.is_finished_length(): - finish_status = FinishStatus(FinishStatus.NO_FINISH) - yield sub_req_id, request_output, metadata, finish_status - else: - continue - else: + now = time.monotonic() + if now >= next_disconnect_check: + next_disconnect_check = now + 1.0 + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, + reason="fetch_pd_stream decode period check network disconnected", + ) + token_list = req_status.pop_all_tokens() + for sub_req_id, request_output, metadata, finish_status in token_list: + output_index = metadata.get("count_output_tokens") + # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 + if output_index == 1: + if first_token_gen is False: + first_token_gen = True + node_run_mode = metadata.pop("node_mode", None) + if node_run_mode == "prefill": + if old_max_new_tokens != 1 and finish_status.is_finished_length(): + finish_status = FinishStatus(FinishStatus.NO_FINISH) yield sub_req_id, request_output, metadata, finish_status + else: + continue + else: + yield sub_req_id, request_output, metadata, finish_status return @@ -471,11 +484,6 @@ async def _wait_to_token_package( 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 ): - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, reason="_wait_to_token_package check network disconnected" - ) - prompt_tokens = metadata["prompt_tokens"] out_token_counter += 1 prompt_cache_len = max(prompt_cache_len, metadata.get("prompt_cache_len", 0)) @@ -578,27 +586,29 @@ async def handle_loop(self): try: for obj in objs: - if obj[0] == ObjType.TOKEN_PACKS: + if obj[0] in (ObjType.TOKEN_PACKS, ObjType.TOKEN_PACKS_COMPACT): token_list, node_load_info = obj[1], obj[2] self.pd_manager.update_node_load_info(node_load_info) - for sub_req_id, text, metadata, finish_status in token_list: - finish_status: FinishStatus = finish_status + compact_pack = obj[0] == ObjType.TOKEN_PACKS_COMPACT + for token_info in token_list: + if compact_pack: + sub_req_id, text, metadata, finish_status_value = unpack_pd_compact_token_info( + token_info + ) + finish_status = FinishStatus(finish_status_value) + else: + sub_req_id, text, metadata, finish_status = token_info group_req_id = convert_sub_id_to_group_id(sub_req_id) - try: - req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.out_token_info_list.append((sub_req_id, text, metadata, finish_status)) - req_status.event.set() - except: - pass + req_status: ReqStatus = self.req_id_to_out_inf.get(group_req_id) + if req_status is not None: + req_status.append_token((sub_req_id, text, metadata, finish_status)) elif obj[0] == ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS: _, group_req_id, prompt_ids = obj try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - async with req_status.lock: - req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids - req_status.prefill_prompt_ids_event.set() + req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids + req_status.prefill_prompt_ids_event.set() except: logger.error( f"PD_UPLOAD_PREFILL_PROMPT_IDS fail find req status for group_req_id: {group_req_id}" @@ -621,7 +631,6 @@ def _split_max_new_tokens(self, max_new_tokens: int) -> List[int]: class ReqStatus: def __init__(self, req_id, p_node, d_node) -> None: self.req_id = req_id - self.lock = asyncio.Lock() self.event = asyncio.Event() self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() @@ -635,19 +644,18 @@ async def wait_to_ready(self): except asyncio.TimeoutError: pass - async def can_read(self, req_id_to_out_inf): - async with self.lock: - self.event.clear() - assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}" - if len(self.out_token_info_list) == 0: - return False - else: - return True - - async def pop_all_tokens(self): - async with self.lock: - ans = self.out_token_info_list.copy() - self.out_token_info_list.clear() + def append_token(self, token_info: Tuple[int, str, dict, FinishStatus]): + # TOKEN_PACKS handling and fetch_pd_stream run on the same event loop. Keeping + # the mutation free of awaits makes the empty -> ready transition atomic. + was_empty = not self.out_token_info_list + self.out_token_info_list.append(token_info) + if was_empty: + self.event.set() + + def pop_all_tokens(self): + self.event.clear() + ans = self.out_token_info_list + self.out_token_info_list = [] return ans diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index ee8892314b..aa97c4902e 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -2,7 +2,7 @@ import time import copy from dataclasses import dataclass, field -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Tuple from lightllm.server.req_id_generator import convert_sub_id_to_group_id from fastapi import WebSocket @@ -41,6 +41,90 @@ class ObjType(enum.Enum): PD_UPLOAD_PREFILL_PROMPT_IDS = 4 # prefill 节点上报生成的 prompt ids 信息。 PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 + TOKEN_PACKS_COMPACT = 7 # 不含 logprobs 等可选字段的紧凑 token 包。 + + +PD_COMPACT_TOKEN_INFO_LEN = 9 +PDCompactTokenInfo = Tuple[ + int, # sub request id + str, # decoded text + int, # count_output_tokens + int, # prompt_tokens + int, # prompt_cache_len + int, # mtp_accepted_token_num + int, # finish status + Optional[str], # node_mode, first token only + Optional[Tuple[int, int, int]], # input text/audio/image tokens, first token only +] +_PD_COMPACT_METADATA_KEYS = frozenset( + { + "count_output_tokens", + "prompt_tokens", + "prompt_cache_len", + "mtp_accepted_token_num", + "node_mode", + "input_usage", + } +) +_PD_INPUT_USAGE_KEYS = frozenset({"input_text_tokens", "input_audio_tokens", "input_image_tokens"}) + + +def build_pd_compact_token_info(sub_req_id, text, metadata, finish_status) -> Optional[PDCompactTokenInfo]: + """Build the lossless compact form, or return None for optional metadata.""" + if not metadata.keys() <= _PD_COMPACT_METADATA_KEYS: + return None + + input_usage = metadata.get("input_usage") + compact_input_usage = None + if input_usage is not None: + if input_usage.keys() != _PD_INPUT_USAGE_KEYS: + return None + compact_input_usage = ( + input_usage["input_text_tokens"], + input_usage["input_audio_tokens"], + input_usage["input_image_tokens"], + ) + + return ( + sub_req_id, + text, + metadata["count_output_tokens"], + metadata["prompt_tokens"], + metadata["prompt_cache_len"], + metadata["mtp_accepted_token_num"], + finish_status.status, + metadata.get("node_mode"), + compact_input_usage, + ) + + +def unpack_pd_compact_token_info(token_info: PDCompactTokenInfo): + ( + sub_req_id, + text, + count_output_tokens, + prompt_tokens, + prompt_cache_len, + mtp_accepted_token_num, + finish_status, + node_mode, + input_usage, + ) = token_info + metadata = { + "count_output_tokens": count_output_tokens, + "prompt_tokens": prompt_tokens, + "prompt_cache_len": prompt_cache_len, + "mtp_accepted_token_num": mtp_accepted_token_num, + } + if node_mode is not None: + metadata["node_mode"] = node_mode + if input_usage is not None: + metadata["input_usage"] = { + "input_text_tokens": input_usage[0], + "input_audio_tokens": input_usage[1], + "input_image_tokens": input_usage[2], + } + return sub_req_id, text, metadata, finish_status @dataclass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 8606ddaabc..9262c76eb3 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -297,7 +297,8 @@ def accept_peer_task_loop( local_trans_task.prefill_page_reg_desc = remote_trans_task.prefill_page_reg_desc local_trans_task.transfer_nbytes = remote_trans_task.transfer_nbytes self.request_page_task_queue.put(local_trans_task) - logger.info(f"recv WRITE request from prefill: {remote_trans_task.to_str()}") + if self.args.detail_log: + logger.info(f"recv WRITE request from prefill: {remote_trans_task.to_str()}") else: # This does not necessarily mean the WRITE protocol state is corrupted. # A common benign case is: decode has already received an abort for this @@ -324,7 +325,8 @@ def accept_peer_task_loop( local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob self.ready_page_task_queue.put(local_trans_task) - logger.info(f"recv WRITE done from prefill: {remote_trans_task.to_str()}") + if self.args.detail_log: + logger.info(f"recv WRITE done from prefill: {remote_trans_task.to_str()}") else: # Same race as the WRITE request stage: decode may have cleaned the # waiting task because the request was aborted, then a late done notify @@ -427,13 +429,14 @@ def success_loop(self): ret = trans_task.createRetObj() self.task_out_queue.put(ret) - if trans_task.start_trans_time is not None: - logger.info( - f"trans task ret success:{ret} cost time: {trans_task.transfer_time()} s " - f"read_page_gpu_time: {read_page_gpu_time_ms:.3f} ms" - ) - else: - logger.info(f"trans task ret success:{ret}") + if self.args.detail_log: + if trans_task.start_trans_time is not None: + logger.info( + f"trans task ret success:{ret} cost time: {trans_task.transfer_time()} s " + f"read_page_gpu_time: {read_page_gpu_time_ms:.3f} ms" + ) + else: + logger.info(f"trans task ret success:{ret}") @log_exception def fail_loop(self): From 7e24cee06976d0ae6da1acfa8a2f3b6fa2632583 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 19 Aug 2026 05:45:29 +0000 Subject: [PATCH 113/214] avoid decode-only allocations on prefill nodes --- lightllm/common/basemodel/basemodel.py | 4 +- lightllm/distributed/communication_op.py | 113 +++++++++++++++++++---- lightllm/models/deepseek_v4/model.py | 3 +- 3 files changed, 100 insertions(+), 20 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 885a02f8dc..e75c0591a2 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -277,7 +277,7 @@ def _init_att_backend1(self): def _init_cudagraph(self): self.graph = ( None - if self.disable_cudagraph + if self.args.run_mode == "prefill" or self.disable_cudagraph else CudaGraph( max_batch_size=self.graph_max_batch_size, max_len_in_batch=self.graph_max_len_in_batch, @@ -318,7 +318,7 @@ def _full_att_decode_autotune(self): Candidate batch sizes follow the same schedule as CUDA Graph capture. Actual benchmarking is delegated to ``fa3_decode_autotune`` in ``sgl_utils``. """ - if self.disable_cudagraph: + if self.args.run_mode == "prefill" or self.disable_cudagraph: return # Only tune on the main model; MTP draft models skip this path. if getattr(self, "is_mtp_draft_model", False): diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 1838828cb1..e1b7a63f4d 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -43,6 +43,56 @@ logger = init_logger(__name__) +def get_deep_ep_prefill_moe_workspace_size( + num_max_tokens_per_rank: int, + hidden_size: int, + intermediate_size: int, + num_experts_per_tok: int, + num_experts: int, + world_size: int, + hidden_dtype: torch.dtype, +) -> int: + """Size one Prefill microbatch workspace for the bounded grouped-MoE path. + + The target chunk covers balanced expert routing in one pass. More skewed + routing remains correct because the consumer already splits expanded rows. + """ + tensor_alignment = 256 + expert_alignment = 128 + metadata_row_granularity = 1024 + + assert num_experts % world_size == 0 + assert intermediate_size % expert_alignment == 0 + + def align_up(value: int, alignment: int) -> int: + return (value + alignment - 1) // alignment * alignment + + hidden_bytes = torch.empty((), dtype=hidden_dtype).element_size() + max_gather_rows = align_up(world_size * num_max_tokens_per_rank, metadata_row_granularity) + num_local_experts = num_experts // world_size + target_chunk_rows = align_up( + num_max_tokens_per_rank * num_experts_per_tok + num_local_experts * (expert_alignment - 1), + expert_alignment, + ) + + gather_out = align_up(max_gather_rows * hidden_size * hidden_bytes, tensor_alignment) + silu_out = align_up(target_chunk_rows * intermediate_size * hidden_bytes, tensor_alignment) + gemm_out_a = align_up(target_chunk_rows * 2 * intermediate_size * hidden_bytes, tensor_alignment) + quant_out = align_up(target_chunk_rows * intermediate_size, tensor_alignment) + quant_scale = align_up(target_chunk_rows * (intermediate_size // expert_alignment) * 4, tensor_alignment) + gemm_out_b = align_up(target_chunk_rows * hidden_size * hidden_bytes, tensor_alignment) + + # TensorBufferManager uses first-fit allocation. W1 keeps silu_out and gemm_out_a + # together; after gemm_out_a is reused by quantization, the freed silu_out block + # can only hold gemm_out_b when it is large enough. + w1_peak = gather_out + silu_out + gemm_out_a + w2_peak = gather_out + silu_out + quant_out + quant_scale + if gemm_out_b > silu_out: + w2_peak += gemm_out_b + # TensorBufferManager may trim an unaligned prefix from the supplied view. + return max(w1_peak, w2_peak) + tensor_alignment + + try: import deep_ep @@ -109,6 +159,7 @@ def __init__(self): self.groups = [] self.ep_buffer = None self.ep_low_latency_buffer = None + self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None @@ -171,12 +222,16 @@ def new_deepep_group( DeepEP legacy low-latency 路径。这里只为实际存在的执行路径分配 buffer, 避免为未使用的路径长期占用显存。 """ - enable_ep_moe = get_env_start_args().enable_ep_moe + args = get_env_start_args() + enable_ep_moe = args.enable_ep_moe prefill_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_prefill() - decode_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() + decode_num_max_dispatch_tokens_per_rank = ( + None if args.run_mode == "prefill" else get_deepep_num_max_dispatch_tokens_per_rank_decode() + ) if not enable_ep_moe: self.ep_buffer = None self.ep_low_latency_buffer = None + self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None return @@ -205,6 +260,7 @@ def new_deepep_group( ) self.ep_mega_moe_buffer = None self.ep_low_latency_buffer = None + self.ep_prefill_moe_workspace = None if not expert_quant_method_names: raise ValueError("No valid MoE quant method was found while initializing DeepEP buffers") @@ -225,10 +281,12 @@ def new_deepep_group( method_name != mega_moe_quant_method for method_name in expert_quant_method_names ) enable_mega_moe_buffer = has_mega_moe_layer - enable_low_latency_buffer = has_legacy_moe_layer else: enable_mega_moe_buffer = False - enable_low_latency_buffer = True + has_legacy_moe_layer = True + + enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" + enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" if enable_low_latency_buffer: # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 @@ -242,6 +300,25 @@ def new_deepep_group( low_latency_mode=True, num_qps_per_rank=(self.ll_num_experts // global_world_size), ) + self.ep_prefill_moe_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( + torch.uint8, use_rdma_buffer=True + ) + + if enable_prefill_workspace: + workspace_size = get_deep_ep_prefill_moe_workspace_size( + num_max_tokens_per_rank=self.ll_num_tokens, + hidden_size=self.ll_hidden, + intermediate_size=moe_intermediate_size, + num_experts_per_tok=num_experts_per_tok, + num_experts=self.ll_num_experts, + world_size=global_world_size, + hidden_dtype=get_torch_dtype(args.data_type), + ) + self.ep_prefill_moe_workspace = torch.empty( + workspace_size * len(self.groups), + dtype=torch.uint8, + device=torch.device("cuda", torch.cuda.current_device()), + ) if enable_mega_moe_buffer: # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy @@ -260,8 +337,10 @@ def new_deepep_group( moe_intermediate_size, ) logger.info( - "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, expert_quant_method_names=%s", + "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " + "mega_moe=%s, expert_quant_method_names=%s", enable_low_latency_buffer, + self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, sorted(expert_quant_method_names), ) @@ -285,15 +364,14 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): logger.warning(f"set num sms for deep_gemm failed: {e}") def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: - """Return a slice of the workspace reused by DeepEP prefill MoE kernels. + """Return a slice of the workspace used by DeepEP Prefill MoE kernels. - DeepEP's low-latency RDMA buffer is idle during prefill, so its local - storage is reused as temporary workspace for the expanded MoE compute - path to reduce peak GPU memory. With one communication group, the - default ``microbatch_index=0`` receives the whole workspace. With - multiple groups, the workspace is split into ``len(self.groups)`` - equal slices and each in-flight microbatch uses the slice matching its - group index. + Pure Prefill nodes own a dedicated workspace and do not initialize a + low-latency Decode buffer. Decode-capable nodes reuse the idle local + RDMA storage. With one communication group, the default + ``microbatch_index=0`` receives the whole workspace. With multiple + groups, the workspace is split into ``len(self.groups)`` equal slices + and each in-flight microbatch uses the slice matching its group index. Args: microbatch_index: Zero-based microbatch and communication-group @@ -307,8 +385,8 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. initialized. The same returned slice must not be used concurrently by overlapping calls. """ - assert self.ep_low_latency_buffer is not None, "DeepEP low-latency buffer is not initialized" - workspace = self.ep_low_latency_buffer.get_local_buffer_tensor(torch.uint8, use_rdma_buffer=True) + assert self.ep_prefill_moe_workspace is not None, "DeepEP Prefill MoE workspace is not initialized" + workspace = self.ep_prefill_moe_workspace microbatch_count = len(self.groups) assert 0 <= microbatch_index < microbatch_count workspace_size = workspace.numel() // microbatch_count @@ -316,8 +394,9 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. def clear_deepep_buffer(self): """ - Prefill MoE compute reuses the low-latency RDMA buffer as workspace. - Clean it before the buffer is used by low-latency decode kernels. + Decode-capable modes reuse the low-latency RDMA buffer during Prefill, + so clean it before the next low-latency Decode. Pure Prefill owns a + dedicated workspace and has no low-latency buffer to clean. """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 0da693f6bd..1035cd9256 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -93,6 +93,7 @@ def _get_compress_rates(self, layer_num): def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() + state_mtp_step = 0 if self.args.run_mode == "prefill" else self.args.mtp_step self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, @@ -102,7 +103,7 @@ def _init_mem_manager(self): compress_rates=self._get_compress_rates(layer_num), indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, - mtp_step=self.args.mtp_step, + mtp_step=state_mtp_step, cpu_cache_token_page_size=( DSV4_CPU_CACHE_TOKEN_PAGE_SIZE if self.args.cpu_cache_token_page_size is None From c211f485aa0c5dcd392a22104f5b37cf44056c2e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 19 Aug 2026 07:44:28 +0000 Subject: [PATCH 114/214] fix: align per-DP request limits with DeepEP capacity --- lightllm/server/api_start.py | 11 ++++++++ lightllm/server/core/objs/start_args_type.py | 1 + .../server/router/req_queue/base_queue.py | 4 ++- lightllm/utils/envs_utils.py | 26 +++++++++++++++++-- 4 files changed, 39 insertions(+), 3 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 07ce95fd2c..f9734e9426 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -115,6 +115,17 @@ def _launch_subprocesses(args: StartArgs): f"graph_max_batch_size to 32" ) + dp_size_in_node = max(1, args.dp // args.nnodes) + args.per_dp_running_max_req_size = args.running_max_req_size // dp_size_in_node + args.graph_max_batch_size = min(args.graph_max_batch_size, args.per_dp_running_max_req_size) + logger.info( + "set per-DP running request limit: global=%d, local_dp=%d, per_dp=%d, graph_max_batch_size=%d", + args.running_max_req_size, + dp_size_in_node, + args.per_dp_running_max_req_size, + args.graph_max_batch_size, + ) + if not args.disable_shm_warning: check_recommended_shm_size(args) diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index ea04be5693..7b75e164b2 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -81,6 +81,7 @@ class StartArgs: ) chat_template: Optional[str] = field(default=None) running_max_req_size: int = field(default=256) + per_dp_running_max_req_size: Optional[int] = field(default=None, init=False) tp: int = field(default=1) dp: int = field(default=1) nnodes: int = field(default=1) diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 9af7afd1b4..b3817472b8 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -25,7 +25,9 @@ def __init__(self, args: StartArgs, router, dp_index, dp_size_in_node) -> None: self.max_total_tokens = args.max_total_token_num - get_fixed_kv_len() assert args.batch_max_tokens is not None self.batch_max_tokens = args.batch_max_tokens - self.running_max_req_size = args.running_max_req_size # Maximum number of concurrent requests + if args.per_dp_running_max_req_size is None: + raise RuntimeError("per_dp_running_max_req_size is not initialized") + self.running_max_req_size = args.per_dp_running_max_req_size self.waiting_req_list: List[Req] = [] # List of queued requests self.router_token_ratio = args.router_token_ratio # ratio to determine whether the router is busy diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 595879f880..e121e2f193 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -38,7 +38,19 @@ def set_env_start_args(args): args = vars(args) os.environ["LIGHTLLM_START_ARGS"] = json.dumps(args) if args["enable_ep_moe"]: - decode_capacity = get_deepep_num_max_dispatch_tokens_per_rank_decode() + if args["run_mode"] == "prefill": + decode_capacity = args["running_max_req_size"] * (args["mtp_step"] + 1) + decode_capacity = ((decode_capacity + 7) // 8) * 8 + configured_decode_capacity = int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", decode_capacity)) + if configured_decode_capacity != decode_capacity: + logger.warning( + "NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE=%d differs from the automatically derived value %d.", + configured_decode_capacity, + decode_capacity, + ) + decode_capacity = max(configured_decode_capacity, decode_capacity) + else: + decode_capacity = get_deepep_num_max_dispatch_tokens_per_rank_decode() min_qp_depth = 2 * (decode_capacity + 1) derived_qp_depth = 1 << (min_qp_depth - 1).bit_length() configured_qp_depth = int(os.getenv("NVSHMEM_QP_DEPTH", derived_qp_depth)) @@ -96,7 +108,17 @@ def get_deepep_num_max_dispatch_tokens_per_rank_prefill(): @lru_cache(maxsize=None) def get_deepep_num_max_dispatch_tokens_per_rank_decode(): args = get_env_start_args() - required = args.running_max_req_size * (args.mtp_step + 1) + per_dp_running_max_req_size = getattr(args, "per_dp_running_max_req_size", None) + if per_dp_running_max_req_size is None: + per_dp_running_max_req_size = args.running_max_req_size + + graph_max_batch_size = 0 + if not args.disable_cudagraph: + graph_max_batch_size = args.graph_max_batch_size + if args.enable_decode_microbatch_overlap: + graph_max_batch_size //= 2 + + required = max(per_dp_running_max_req_size, graph_max_batch_size) * (args.mtp_step + 1) required = ((required + 7) // 8) * 8 configured = int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", required)) if configured != required: From bc347a9110bd532556b1eddac3d4528d666ae2bb Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 20 Aug 2026 02:21:35 +0000 Subject: [PATCH 115/214] fix ttft --- lightllm/server/httpserver_for_pd_master/manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index d62ab85871..19a930f214 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -248,7 +248,7 @@ async def _generate_one( results_generator = self._wait_to_token_package( p_node, d_node, - start_time, + start_time if iter_index == 0 else time.time(), block_prompt, sampling_params, multimodal_params, From a3160f25b30bdad7856e0e77b7daff1254de9717 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 22 Jul 2026 15:49:26 +0800 Subject: [PATCH 116/214] feat: imbalance statistics --- lightllm/common/basemodel/basemodel.py | 13 +- .../meta_weights/fused_moe/ep_balance.py | 14 + .../fused_moe/impl/deepgemm_impl.py | 24 +- .../fused_moe/grouped_fused_moe_ep.py | 10 + lightllm/distributed/communication_op.py | 9 + lightllm/server/api_cli.py | 5 + lightllm/server/core/objs/start_args_type.py | 1 + lightllm/server/metrics/metrics.py | 12 + .../mode_backend/ep_balance_monitor.py | 346 +++++++++++ .../server/router/model_infer/model_rpc.py | 9 + .../model_infer/test_ep_balance_monitor.py | 546 ++++++++++++++++++ 11 files changed, 984 insertions(+), 5 deletions(-) create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py create mode 100644 lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py create mode 100644 unit_tests/server/router/model_infer/test_ep_balance_monitor.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index e75c0591a2..f025fb9d21 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -65,6 +65,7 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() + self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] self.max_total_token_num = kvargs["max_total_token_num"] @@ -361,9 +362,14 @@ def forward(self, model_input: ModelInput): assert model_input.mem_indexes.is_cuda if model_input.is_prefill: - return self._prefill(model_input) - else: - return self._decode(model_input) + model_output = self._prefill(model_input) + self._record_prefill_ep_balance() + return model_output + return self._decode(model_input) + + def _record_prefill_ep_balance(self): + if self.ep_balance_monitor is not None: + self.ep_balance_monitor.record_prefill_round() def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -865,6 +871,7 @@ def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: Mod dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event + self._record_prefill_ep_balance() return model_output0, model_output1 @torch.no_grad() diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py new file mode 100644 index 0000000000..70435146c5 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py @@ -0,0 +1,14 @@ +from dataclasses import dataclass + + +@dataclass(slots=True) +class PrefillEPBalanceCounters: + """Cumulative CPU loads for one EP MoE layer's completed prefill dispatches.""" + + route_load: int = 0 + compute_load: int = 0 + + def accumulate(self, route_load: int, compute_load: int): + """Accumulate exact route and alignment-expanded compute loads for one prefill dispatch.""" + self.route_load += route_load + self.compute_load += compute_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 103cbc3087..417910f89c 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -20,6 +20,10 @@ class FuseMoeDeepGEMM(FuseMoeTriton): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.ep_balance_counters = None + def _select_experts( self, input_tensor: torch.Tensor, @@ -91,6 +95,7 @@ def _fused_experts( previous_event=None, # for overlap clamp_limit=clamp_limit, alloc_tensor_func=alloc_tensor_func, + ep_balance_counters=self.ep_balance_counters, ) return output @@ -200,8 +205,23 @@ def dispatch( use_tma_aligned_col_major_sf=True, ) - def hook(): - event.current_stream_wait() + counters = self.ep_balance_counters + if counters is None: + + def hook(): + event.current_stream_wait() + + else: + # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. + route_load = topk_idx.numel() + compute_load = recv_x[0].shape[0] + + def hook(): + event.current_stream_wait() + counters.accumulate( + route_load=route_load, + compute_load=compute_load, + ) return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 04b54ea420..210da61da1 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -19,6 +19,7 @@ ep_gather_chunk, ep_zero_padding, ) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -210,6 +211,7 @@ def fused_experts( previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, alloc_tensor_func: Callable = torch.empty, + ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): @@ -242,6 +244,7 @@ def fused_experts( previous_event=previous_event, clamp_limit=clamp_limit, alloc_tensor_func=alloc_tensor_func, + ep_balance_counters=ep_balance_counters, ) @@ -262,6 +265,7 @@ def fused_experts_impl( previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, alloc_tensor_func: Callable = torch.empty, + ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -314,6 +318,12 @@ def fused_experts_impl( do_expand=True, use_tma_aligned_col_major_sf=True, ) + if ep_balance_counters is not None: + # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. + ep_balance_counters.accumulate( + route_load=topk_idx.numel(), + compute_load=recv_x[0].shape[0], + ) # Dispatch is synchronous in this path. Its FP8 source is no longer # needed once the received tensors have been produced. del qinput_tensor, input_scale diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index e1b7a63f4d..330db23bb1 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -157,6 +157,7 @@ def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, as class DistributeGroupManager: def __init__(self): self.groups = [] + self.ep_balance_monitor_group = None self.ep_buffer = None self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None @@ -175,6 +176,14 @@ def create_groups(self, group_size: int): if not args.disable_flashinfer_allreduce: group.init_flashinfer_reduce() self.groups.append(group) + if ( + getattr(args, "enable_ep_moe", False) + and not getattr(args, "disable_ep_balance_monitor", False) + and getattr(args, "run_mode", "normal") != "decode" + and not getattr(args, "enable_prefill_cudagraph", False) + and not is_sm100_gpu() + ): + self.ep_balance_monitor_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") return def get_default_group(self) -> CustomProcessGroup: diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 2523f78b81..0206bd863b 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -719,6 +719,11 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Whether to enable ep moe for deepseekv3 model.""", ) + parser.add_argument( + "--disable_ep_balance_monitor", + action="store_true", + help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", + ) parser.add_argument( "--ep_redundancy_expert_config_path", type=str, diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 7b75e164b2..ecfb111f26 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -179,6 +179,7 @@ class StartArgs: default="gpu_counter", metadata={"choices": ["cpu_counter", "pin_mem_counter", "gpu_counter"]} ) enable_ep_moe: bool = field(default=False) + disable_ep_balance_monitor: bool = field(default=False) ep_redundancy_expert_config_path: Optional[str] = field(default=None) auto_update_redundancy_expert: bool = field(default=False) enable_fused_shared_experts: bool = field(default=False) diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index 0d42462c3f..c19d756c17 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -32,6 +32,15 @@ "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": ( + "Estimated critical-path excess compute per logical routed source token, GFLOPs/token" + ), + "lightllm_prefill_ep_compute_critical_overhead_ratio": ( + "Estimated excess critical compute divided by balanced compute; 0.3 means +30%" + ), + "lightllm_prefill_ep_placement_pressure_drift": ( + "Normalized temporal drift of overloaded-rank pressure from the latest complete prefill report" + ), } @@ -111,6 +120,9 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") + self.create_gauge("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token") + self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio") + self.create_gauge("lightllm_prefill_ep_placement_pressure_drift") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py new file mode 100644 index 0000000000..107f4ba504 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py @@ -0,0 +1,346 @@ +import threading +from array import array +from typing import Optional, Tuple + +import torch +import torch.distributed as dist + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters +from lightllm.distributed.communication_op import dist_group_manager +from lightllm.server.metrics.manager import MetricClient +from lightllm.utils.device_utils import is_sm100_gpu +from lightllm.utils.dist_utils import ( + get_global_rank, + get_global_world_size, +) +from lightllm.utils.log_utils import init_logger +from lightllm.utils.shm_port_args import get_shm_port_args + + +logger = init_logger(__name__) + +EP_BALANCE_PREFILL_ROUNDS_PER_REPORT = 100 +EP_BALANCE_ROUND_BUFFER_CAPACITY = 4096 +EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS = 20 +EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD = 0.10 +ROUTE_LOAD = 0 +COMPUTE_LOAD = 1 +GFLOP = 1_000_000_000 + + +def should_enable_ep_balance_monitor(args) -> bool: + if args.enable_prefill_cudagraph or is_sm100_gpu(): + return False + return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" + + +def calculate_prefill_balance_stats( + round_stats: torch.Tensor, # [num_rounds, num_layers, world_size, 2] (route/compute) + layer_routed_experts: torch.Tensor, # [num_layers] + layer_flops_per_expert_token: torch.Tensor, # [num_layers] + layer_topks: torch.Tensor, # [num_layers] + source_token_replication: int, + report_min_route_samples_per_expert: int = 100, +) -> Optional[dict]: + """Summarize complete-prefill samples from [round, layer, rank, route/compute]. + + MoE layers execute sequentially, and every layer waits for its slowest EP + rank. Preserve the layer dimension until after taking the cross-rank max so + that different slow ranks in different layers cannot cancel each other. + """ + assert source_token_replication > 0 + + layer_route_load = round_stats[:, :, :, ROUTE_LOAD].sum(dim=(0, 2)) + minimum_route_samples = layer_routed_experts * report_min_route_samples_per_expert + if torch.any(layer_route_load < minimum_route_samples): + return None + total_route_load = layer_route_load.sum() + + compute_rank_load = round_stats[:, :, :, COMPUTE_LOAD] + compute_rank_load_float = compute_rank_load.to(torch.float64) + excess_compute_load = compute_rank_load_float.max(dim=2).values - compute_rank_load_float.mean(dim=2) + + # Every MoE expert token executes the two projections packed in w13 plus + # the w2 projection. Weight each padded compute token by the layer's + # actual matrix sizes so the metric remains comparable across models. + excess_compute_flops = (excess_compute_load * layer_flops_per_expert_token.to(torch.float64)).sum() + balanced_compute_flops = ( + compute_rank_load_float.mean(dim=2) * layer_flops_per_expert_token.to(torch.float64) + ).sum() + if balanced_compute_flops == 0: + return None + + # Non-TPSP prefill gathers one route-load copy per TP rank. Divide the + # replica count out so GFLOP/token uses logical source tokens. + source_tokens = total_route_load.to(torch.float64) / ( + layer_topks.to(torch.float64).sum() * source_token_replication + ) + if source_tokens == 0: + return None + + return { + "prefill_rounds": int(compute_rank_load.shape[0]), + "critical_overhead_gflops_per_routed_token": float((excess_compute_flops / source_tokens / GFLOP).item()), + "prefill_ep_compute_critical_overhead_ratio": float((excess_compute_flops / balanced_compute_flops).item()), + } + + +def calculate_prefill_placement_pressure_drift( + round_stats: torch.Tensor, + previous_pressure_signature: Optional[torch.Tensor] = None, + bucket_rounds: int = EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, +) -> Tuple[float, torch.Tensor]: + """Measure how prefill rank-pressure placement changes across time buckets. + + This is rank-0 CPU-only report analysis. The returned final bucket is a + compact signature that allows the next report to include the boundary pair. + """ + if bucket_rounds <= 0: + raise ValueError(f"bucket_rounds must be positive, got {bucket_rounds}") + if round_stats.ndim != 4 or round_stats.shape[-1] != 2: + raise ValueError( + "round_stats must have shape [num_rounds, num_layers, world_size, 2], " f"got {tuple(round_stats.shape)}" + ) + num_rounds, num_layers, world_size, _ = round_stats.shape + if num_rounds <= 0: + raise ValueError("round_stats must contain at least one round") + if num_layers <= 0 or world_size <= 0: + raise ValueError("round_stats must contain at least one layer and rank") + if num_rounds % bucket_rounds != 0: + raise ValueError(f"num_rounds ({num_rounds}) must be divisible by bucket_rounds ({bucket_rounds})") + + num_buckets = num_rounds // bucket_rounds + bucket_rank_load = ( + round_stats[:, :, :, COMPUTE_LOAD] + .to(torch.float64) + .reshape(num_buckets, bucket_rounds, num_layers, world_size) + .sum(dim=1) + ) + mean_rank_load = bucket_rank_load.mean(dim=2, keepdim=True).clamp_min(1) + pressure = torch.relu(bucket_rank_load / mean_rank_load - 1) + + if previous_pressure_signature is not None: + expected_shape = (num_layers, world_size) + if tuple(previous_pressure_signature.shape) != expected_shape: + raise ValueError( + "previous_pressure_signature must have shape " + f"{expected_shape}, got {tuple(previous_pressure_signature.shape)}" + ) + left = torch.cat((previous_pressure_signature.to(torch.float64).unsqueeze(0), pressure[:-1]), dim=0) + right = pressure + else: + left = pressure[:-1] + right = pressure[1:] + + total_pressure = (left + right).sum() + if total_pressure == 0: + drift = 0.0 + else: + drift = float((left - right).abs().sum().div(total_pressure).item()) + return drift, pressure[-1].clone() + + +def classify_prefill_placement_pressure_drift(drift: float) -> str: + if drift < EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD: + return "stable" + return "dynamic" + + +def _find_fused_moe_weights(model): + weights_by_id = {} + for layer in model.trans_layers_weight: + for value in getattr(layer, "__dict__", {}).values(): + if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: + weights_by_id[id(value)] = value + return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) + + +class EPBalanceMonitor: + """Report cross-rank imbalance for non-overlapping blocks of complete prefill rounds.""" + + def __init__(self, model: TpPartBaseModel): + self.global_rank = get_global_rank() + self.world_size = get_global_world_size() + self.weights = _find_fused_moe_weights(model) + self.enabled = bool(self.weights) + + if not self.enabled: + return + + self.source_token_replication = 1 if model.args.enable_tpsp_mix_mode else model.tp_world_size_ + self.counters: list[PrefillEPBalanceCounters] = [PrefillEPBalanceCounters() for _ in self.weights] + for weight, counter in zip(self.weights, self.counters): + weight.fuse_moe_impl.ep_balance_counters = counter + self.layer_routed_experts = torch.tensor( + [weight.n_routed_experts for weight in self.weights], dtype=torch.int64 + ) + self.layer_flops_per_expert_token = torch.tensor( + [ + # Each expert-token performs gate, up, and down projections; each MAC counts as 2 FLOPs. + 2 * 3 * weight.hidden_size * weight.moe_intermediate_size + for weight in self.weights + ], + dtype=torch.float64, + ) + self.layer_topks = torch.tensor( + [weight.num_experts_per_tok for weight in self.weights], + dtype=torch.float64, + ) + self._round_buffer_storage = array("q", [0]) * (EP_BALANCE_ROUND_BUFFER_CAPACITY * len(self.weights) * 2) + self._round_buffer = torch.frombuffer(self._round_buffer_storage, dtype=torch.int64).view( + EP_BALANCE_ROUND_BUFFER_CAPACITY, len(self.weights), 2 + ) + self._round_ready = threading.Event() + self._written_round_count = 0 # Prefill rounds fully written to the ring buffer. + self._processed_round_count = 0 # Prefill rounds consumed by the monitor thread. + self._overflowed = False + self._previous_pressure_signature: Optional[torch.Tensor] = None + self._common_round_end = torch.zeros((), dtype=torch.int64) + + self.gloo_group = dist_group_manager.ep_balance_monitor_group + if self.gloo_group is None: + raise RuntimeError("EP balance monitor requires a pre-created dedicated Gloo process group") + self.metric_client = MetricClient(get_shm_port_args().metric_port) if self.global_rank == 0 else None + threading.Thread(target=self._monitor_loop, daemon=True, name="ep-balance-monitor").start() + + def record_prefill_round(self): + """Publish one complete all-layer prefill sample to the SPSC ring.""" + if not self.enabled: + return + + written_round_count = self._written_round_count + if written_round_count - self._processed_round_count >= EP_BALANCE_ROUND_BUFFER_CAPACITY: + if not self._overflowed: + self._overflowed = True + self._round_ready.set() + return + + storage_index = (written_round_count % EP_BALANCE_ROUND_BUFFER_CAPACITY) * len(self.counters) * 2 + for counter in self.counters: + self._round_buffer_storage[storage_index] = counter.route_load + self._round_buffer_storage[storage_index + 1] = counter.compute_load + counter.route_load = 0 + counter.compute_load = 0 + storage_index += 2 + + # Publish only after the entire slot is written. The SPSC producer and + # monitor thread run under the CPython GIL, so this count is the release + # point for the corresponding ring slot. + self._written_round_count = written_round_count + 1 + if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + self._round_ready.set() + + def _get_common_round_end(self) -> int: + """Return the exclusive round boundary completed by every rank.""" + self._common_round_end.fill_(self._written_round_count) + dist.all_reduce(self._common_round_end, op=dist.ReduceOp.MIN, group=self.gloo_group) + return int(self._common_round_end.item()) + + def _raise_buffer_overflow(self, phase: str, common_round_end: Optional[int] = None): + message = ( + "EP balance prefill-round buffer overflowed " + f"phase={phase} written={self._written_round_count} " + f"processed={self._processed_round_count} capacity={EP_BALANCE_ROUND_BUFFER_CAPACITY}" + ) + if common_round_end is not None: + message += f" common_round_end={common_round_end}" + raise RuntimeError(message) + + def _copy_local_rounds(self, start: int, end: int) -> torch.Tensor: + """Copy local prefill-round loads in the half-open range [start, end).""" + num_rounds = end - start + if num_rounds > EP_BALANCE_ROUND_BUFFER_CAPACITY: + raise ValueError("requested EP balance round range exceeds ring capacity") + start_index = start % EP_BALANCE_ROUND_BUFFER_CAPACITY + if start_index + num_rounds <= EP_BALANCE_ROUND_BUFFER_CAPACITY: + return self._round_buffer[start_index : start_index + num_rounds].clone() + end_index = (start_index + num_rounds) % EP_BALANCE_ROUND_BUFFER_CAPACITY + return torch.cat((self._round_buffer[start_index:], self._round_buffer[:end_index]), dim=0) + + def _gather_round_stats(self, local_round_stats: torch.Tensor) -> Optional[torch.Tensor]: + """Gather rank-local stats as [round, layer, rank, route/compute].""" + gathered = ( + [torch.empty_like(local_round_stats) for _ in range(self.world_size)] if self.global_rank == 0 else None + ) + dist.gather(local_round_stats, gather_list=gathered, dst=0, group=self.gloo_group) + if self.global_rank != 0: + return None + # [rank, round, layer, route/compute] + # -> [round, layer, rank, route/compute] + return torch.stack(gathered).permute(1, 2, 0, 3) + + def _log_stats(self, round_stats: torch.Tensor): + """Compute and log balance statistics for one complete global window.""" + compute = calculate_prefill_balance_stats( + round_stats, + self.layer_routed_experts, + self.layer_flops_per_expert_token, + self.layer_topks, + self.source_token_replication, + ) + if compute is None: + return + + drift, self._previous_pressure_signature = calculate_prefill_placement_pressure_drift( + round_stats, + previous_pressure_signature=self._previous_pressure_signature, + ) + drift_state = classify_prefill_placement_pressure_drift(drift) + + logger.info( + "ep_balance " + f"phase=prefill prefill_rounds={compute['prefill_rounds']} " + "prefill_ep_critical_overhead_gflops_per_routed_token=" + f"{compute['critical_overhead_gflops_per_routed_token']:.4f} " + "prefill_ep_compute_critical_overhead_ratio=" + f"{compute['prefill_ep_compute_critical_overhead_ratio']:.4f} " + f"prefill_ep_placement_pressure_drift={drift:.4f} " + f"prefill_ep_placement_pressure_state={drift_state}" + ) + self.metric_client.gauge_set( + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", + compute["critical_overhead_gflops_per_routed_token"], + ) + self.metric_client.gauge_set( + "lightllm_prefill_ep_compute_critical_overhead_ratio", + compute["prefill_ep_compute_critical_overhead_ratio"], + ) + self.metric_client.gauge_set("lightllm_prefill_ep_placement_pressure_drift", drift) + + def _monitor_loop(self): + """Consume commonly completed rounds in background report-sized windows.""" + try: + while True: + self._round_ready.wait() + self._round_ready.clear() + if self._overflowed: + self._raise_buffer_overflow("before_sync") + common_round_end = self._get_common_round_end() + if common_round_end - self._processed_round_count > EP_BALANCE_ROUND_BUFFER_CAPACITY: + self._raise_buffer_overflow("common_round_lag", common_round_end=common_round_end) + + while common_round_end - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + round_start = self._processed_round_count + round_end = round_start + EP_BALANCE_PREFILL_ROUNDS_PER_REPORT + local_round_stats = self._copy_local_rounds(round_start, round_end) + if self._overflowed: + self._raise_buffer_overflow("after_copy") + round_stats = self._gather_round_stats(local_round_stats) + self._processed_round_count = round_end + if self.global_rank == 0: + self._log_stats(round_stats) + + if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: + self._round_ready.set() + except Exception as exc: + logger.exception(f"EP balance monitor stopped unexpectedly: {exc}") + self._disable() + return + + def _disable(self): + """Detach counters from MoE weights and disable monitoring.""" + for weight in self.weights: + weight.fuse_moe_impl.ep_balance_counters = None + self.enabled = False diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 1628830f9f..b1d7e50fe5 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -28,6 +28,10 @@ ) from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps +from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( + EPBalanceMonitor, + should_enable_ep_balance_monitor, +) from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.utils.log_utils import init_logger from lightllm.utils.graceful_utils import graceful_registry @@ -108,6 +112,11 @@ def exposed_init_model(self, kvargs): logger.info("init redundancy_expert_manager") else: self.redundancy_expert_manager = None + + if should_enable_ep_balance_monitor(self.args): + monitor = EPBalanceMonitor(self.backend.model) + if monitor.enabled: + self.backend.model.ep_balance_monitor = monitor return def exposed_get_max_total_token_num(self): diff --git a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py new file mode 100644 index 0000000000..2ddcaafddf --- /dev/null +++ b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py @@ -0,0 +1,546 @@ +import threading +from array import array +from types import SimpleNamespace + +import pytest +import torch +from prometheus_client import generate_latest + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters +from lightllm.distributed import communication_op as communication_op_module +from lightllm.server.metrics.metrics import Monitor +from lightllm.server.router.model_infer.mode_backend import ep_balance_monitor as monitor_module +from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( + EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, + calculate_prefill_placement_pressure_drift, + calculate_prefill_balance_stats, + classify_prefill_placement_pressure_drift, + should_enable_ep_balance_monitor, +) + + +@pytest.fixture(autouse=True) +def _mock_non_sm100(monkeypatch): + monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: False) + + +def _stats(source_token_replication: int): + return calculate_prefill_balance_stats( + torch.tensor([[[[800, 40], [800, 20]], [[800, 20], [800, 40]]]], dtype=torch.int64), + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=source_token_replication, + ) + + +def _monitor_args(**overrides): + args = { + "enable_ep_moe": True, + "disable_ep_balance_monitor": False, + "run_mode": "normal", + "enable_prefill_cudagraph": False, + } + args.update(overrides) + return SimpleNamespace(**args) + + +def _pressure_round_stats(bucket_rank_loads): + """Build [round, layer=1, rank, route/compute] CPU samples for drift tests.""" + return torch.tensor([[[[0, load] for load in rank_loads]] for rank_loads in bucket_rank_loads], dtype=torch.int64) + + +def test_pressure_drift_is_zero_for_identical_pressure(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [2, 1]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_is_one_for_complete_hot_rank_migration(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0], [0, 2]]), bucket_rounds=1) + assert drift == 1.0 + + +def test_pressure_drift_tracks_same_rank_magnitude_change(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [3, 1]]), bucket_rounds=1) + assert drift == pytest.approx(0.2) + + +def test_pressure_drift_is_invariant_to_uniform_load_scale(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [4, 2]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_is_zero_for_balanced_inputs(): + drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[8, 8], [16, 16]]), bucket_rounds=1) + assert drift == 0.0 + + +def test_pressure_drift_previous_signature_bridges_report_boundary(): + _, signature = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0]]), bucket_rounds=1) + drift, next_signature = calculate_prefill_placement_pressure_drift( + _pressure_round_stats([[0, 2]]), previous_pressure_signature=signature, bucket_rounds=1 + ) + assert drift == 1.0 + assert torch.equal(next_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) + + +def test_pressure_drift_default_bucket_handles_normal_report_window(): + round_stats = _pressure_round_stats([[4, 2]] * 100) + drift, signature = calculate_prefill_placement_pressure_drift(round_stats) + assert EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS == 20 + assert drift == 0.0 + assert signature.shape == (1, 2) + + +@pytest.mark.parametrize( + ("drift", "expected"), + [ + (0.0, "stable"), + (0.0999, "stable"), + (0.10, "dynamic"), + (0.2999, "dynamic"), + (0.30, "dynamic"), + (0.75, "dynamic"), + (1.0, "dynamic"), + ], +) +def test_pressure_drift_classification_boundaries(drift, expected): + assert classify_prefill_placement_pressure_drift(drift) == expected + + +def test_monitor_log_stats_reports_pressure_drift_and_bridges_reports(monkeypatch): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.layer_routed_experts = torch.tensor([1], dtype=torch.int64) + monitor.layer_flops_per_expert_token = torch.tensor([2.0], dtype=torch.float64) + monitor.layer_topks = torch.tensor([2.0], dtype=torch.float64) + monitor.source_token_replication = 1 + monitor._previous_pressure_signature = None + metric_calls = [] + monitor.metric_client = SimpleNamespace(gauge_set=lambda name, value: metric_calls.append((name, value))) + logs = [] + monkeypatch.setattr(monitor_module.logger, "info", logs.append) + + def report_for_hot_rank(hot_rank): + round_stats = torch.zeros((100, 1, 2, 2), dtype=torch.int64) + round_stats[:, :, :, monitor_module.ROUTE_LOAD] = 100 + round_stats[:, :, hot_rank, monitor_module.COMPUTE_LOAD] = 2 + monitor._log_stats(round_stats) + + report_for_hot_rank(0) + first_signature = monitor._previous_pressure_signature.clone() + report_for_hot_rank(1) + + assert "prefill_ep_placement_pressure_drift=0.0000" in logs[0] + assert "prefill_ep_placement_pressure_state=stable" in logs[0] + assert "prefill_ep_placement_pressure_drift=0.2000" in logs[1] + assert "prefill_ep_placement_pressure_state=dynamic" in logs[1] + assert torch.equal(first_signature, torch.tensor([[1.0, 0.0]], dtype=torch.float64)) + assert torch.equal(monitor._previous_pressure_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) + assert metric_calls == [ + ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), + ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), + ("lightllm_prefill_ep_placement_pressure_drift", 0.0), + ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), + ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), + ("lightllm_prefill_ep_placement_pressure_drift", pytest.approx(0.2)), + ] + + +def test_critical_overhead_preserves_per_layer_slowest_rank(): + stats = _stats(source_token_replication=1) + assert stats is not None + assert stats["critical_overhead_gflops_per_routed_token"] == pytest.approx(7.5e-11) + assert stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx(1 / 3) + + +def test_non_tpsp_tp_replication_scales_gflops_per_token_but_not_ratio(): + tpsp_stats = _stats(source_token_replication=1) + non_tpsp_tp8_stats = _stats(source_token_replication=8) + assert tpsp_stats is not None and non_tpsp_tp8_stats is not None + assert non_tpsp_tp8_stats["critical_overhead_gflops_per_routed_token"] == pytest.approx( + tpsp_stats["critical_overhead_gflops_per_routed_token"] * 8 + ) + assert non_tpsp_tp8_stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx( + tpsp_stats["prefill_ep_compute_critical_overhead_ratio"] + ) + + +def test_critical_overhead_is_zero_when_ranks_are_balanced(): + stats = calculate_prefill_balance_stats( + torch.tensor([[[[100, 32], [100, 32]], [[100, 64], [100, 64]]]], dtype=torch.int64), + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=1, + ) + assert stats is not None + assert stats["critical_overhead_gflops_per_routed_token"] == 0.0 + assert stats["prefill_ep_compute_critical_overhead_ratio"] == 0.0 + + +@pytest.mark.parametrize( + "round_stats", + [ + torch.tensor([[[[1, 1], [1, 1]]]], dtype=torch.int64), + torch.tensor([[[[100, 0], [100, 0]]]], dtype=torch.int64), + ], +) +def test_critical_overhead_rejects_insufficient_or_zero_compute_samples(round_stats): + assert ( + calculate_prefill_balance_stats( + round_stats, + layer_routed_experts=torch.tensor([1]), + layer_flops_per_expert_token=torch.tensor([2.0]), + layer_topks=torch.tensor([2.0]), + source_token_replication=1, + ) + is None + ) + + +def test_cpu_counter_accumulates_multiple_prefill_dispatches(): + counters = PrefillEPBalanceCounters() + counters.accumulate(route_load=3, compute_load=128) + counters.accumulate(route_load=4, compute_load=256) + assert (counters.route_load, counters.compute_load) == (7, 384) + + +def test_monitor_reuses_manager_precreated_dedicated_gloo_group(monkeypatch): + sentinel_group = object() + impl = SimpleNamespace(ep_balance_counters=None) + weight = SimpleNamespace( + fuse_moe_impl=impl, + n_routed_experts=8, + hidden_size=16, + moe_intermediate_size=32, + num_experts_per_tok=2, + ) + model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) + + monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) + monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) + metric_ports = [] + monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=4321)) + monkeypatch.setattr(monitor_module, "MetricClient", lambda port: metric_ports.append(port) or object()) + monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) + monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) + monkeypatch.setattr( + monitor_module.threading, + "Thread", + lambda *args, **kwargs: SimpleNamespace(start=lambda: None), + ) + + monitor = monitor_module.EPBalanceMonitor(model) + + assert monitor.gloo_group is sentinel_group + assert impl.ep_balance_counters is monitor.counters[0] + assert metric_ports == [4321] + + +def test_nonzero_rank_monitor_does_not_create_metric_client(monkeypatch): + sentinel_group = object() + impl = SimpleNamespace(ep_balance_counters=None) + weight = SimpleNamespace( + fuse_moe_impl=impl, + n_routed_experts=8, + hidden_size=16, + moe_intermediate_size=32, + num_experts_per_tok=2, + ) + model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) + metric_client_calls = [] + monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) + monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 1) + monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) + monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) + monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: pytest.fail("unexpected port lookup")) + monkeypatch.setattr( + monitor_module, + "MetricClient", + lambda port: metric_client_calls.append(port) or pytest.fail("unexpected metric client"), + ) + monkeypatch.setattr( + monitor_module.threading, + "Thread", + lambda *args, **kwargs: SimpleNamespace(start=lambda: None), + ) + + monitor = monitor_module.EPBalanceMonitor(model) + + assert monitor.metric_client is None + assert metric_client_calls == [] + + +@pytest.mark.parametrize("disable_monitor", [False, True]) +def test_group_manager_creates_monitor_gloo_group_only_when_enabled(monkeypatch, disable_monitor): + monitor_group = object() + custom_groups = [] + + class FakeCustomProcessGroup: + def init_symm_mem_reduce(self): + pass + + def init_flashinfer_reduce(self): + pass + + args = SimpleNamespace( + enable_ep_moe=True, + disable_ep_balance_monitor=disable_monitor, + run_mode="normal", + enable_prefill_cudagraph=False, + disable_symm_mem_allreduce=True, + disable_flashinfer_allreduce=True, + ) + monkeypatch.setattr(communication_op_module, "get_env_start_args", lambda: args) + monkeypatch.setattr( + communication_op_module, + "CustomProcessGroup", + lambda: custom_groups.append(FakeCustomProcessGroup()) or custom_groups[-1], + ) + monkeypatch.setattr(communication_op_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(communication_op_module, "is_sm100_gpu", lambda: False) + calls = [] + monkeypatch.setattr( + communication_op_module.dist, + "new_group", + lambda *args, **kwargs: calls.append((args, kwargs)) or monitor_group, + ) + + manager = communication_op_module.DistributeGroupManager() + manager.create_groups(group_size=2) + + assert len(manager.groups) == 2 + if disable_monitor: + assert calls == [] + assert manager.ep_balance_monitor_group is None + else: + assert calls == [((), {"ranks": [0, 1], "backend": "gloo"})] + assert manager.ep_balance_monitor_group is monitor_group + + +def test_monitor_registers_prefill_ep_gauges_with_model_label(): + monitor = Monitor( + SimpleNamespace( + metric_gateway=None, + job_name="test", + grouping_key=[], + enable_monitor_auth=False, + model_name="monitor-test-model", + max_req_total_len=128, + mtp_step=0, + ) + ) + values = { + "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": 1.25, + "lightllm_prefill_ep_compute_critical_overhead_ratio": 0.3, + "lightllm_prefill_ep_placement_pressure_drift": 0.125, + } + assert set(values).issubset(monitor.monitor_registry) + for name, value in values.items(): + monitor.gauge_set(name, value) + + exposition = generate_latest(monitor.registry).decode() + for name, value in values.items(): + assert f'{name}{{model_name="monitor-test-model"}} {value}' in exposition + + +def test_record_prefill_round_stores_cumulative_counter_deltas_in_ring_buffer(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters()] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = 0 + monitor._processed_round_count = 0 + monitor._overflowed = False + + monitor.counters[0].accumulate(route_load=3, compute_load=128) + monitor.record_prefill_round() + assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (0, 0) + monitor.counters[0].accumulate(route_load=2, compute_load=256) + monitor.record_prefill_round() + + assert torch.equal( + monitor._copy_local_rounds(0, 2), + torch.tensor([[[3, 128]], [[2, 256]]], dtype=torch.int64), + ) + assert not hasattr(monitor, "_round_lock") + + +def test_spsc_ring_copy_wraps_without_a_lock(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters()] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 + monitor._processed_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 + monitor._overflowed = False + + for value in (11, 12, 13): + monitor.counters[0].accumulate(route_load=value, compute_load=value * 10) + monitor.record_prefill_round() + + assert torch.equal( + monitor._copy_local_rounds(monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2, monitor._written_round_count), + torch.tensor([[[11, 110]], [[12, 120]], [[13, 130]]], dtype=torch.int64), + ) + assert not hasattr(monitor, "_round_lock") + + +def test_spsc_ring_overflow_is_deferred_to_the_monitor_thread(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.enabled = True + monitor.counters = [PrefillEPBalanceCounters(route_load=7, compute_load=70)] + monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) + monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( + monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 + ) + monitor._round_ready = threading.Event() + monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY + monitor._processed_round_count = 0 + monitor._overflowed = False + + monitor.record_prefill_round() + + assert monitor._overflowed + assert monitor._round_ready.is_set() + assert monitor._written_round_count == monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY + assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (7, 70) + + +def test_raise_buffer_overflow_always_reports_phase_and_ring_counts(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor._written_round_count = 23 + monitor._processed_round_count = 7 + + with pytest.raises(RuntimeError) as exc_info: + monitor._raise_buffer_overflow("before_sync") + + assert str(exc_info.value) == ( + "EP balance prefill-round buffer overflowed " + f"phase=before_sync written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY}" + ) + + +def test_raise_buffer_overflow_optionally_reports_common_round_end(): + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor._written_round_count = 23 + monitor._processed_round_count = 7 + + with pytest.raises(RuntimeError) as exc_info: + monitor._raise_buffer_overflow("common_round_lag", common_round_end=19) + + assert str(exc_info.value) == ( + "EP balance prefill-round buffer overflowed " + f"phase=common_round_lag written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY} " + "common_round_end=19" + ) + + +def test_gather_round_stats_only_allocates_receive_buffers_on_rank_zero(monkeypatch): + local_round_stats = torch.tensor([[[3, 128]]], dtype=torch.int64) + sentinel_group = object() + + rank_zero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + rank_zero_monitor.global_rank = 0 + rank_zero_monitor.world_size = 2 + rank_zero_monitor.gloo_group = sentinel_group + + def root_gather(input_tensor, gather_list, dst, group): + assert dst == 0 and group is sentinel_group + assert len(gather_list) == 2 + gather_list[0].copy_(input_tensor) + gather_list[1].copy_(input_tensor + 1) + + monkeypatch.setattr(monitor_module.dist, "gather", root_gather) + result = rank_zero_monitor._gather_round_stats(local_round_stats) + assert torch.equal(result, torch.tensor([[[[3, 128], [4, 129]]]], dtype=torch.int64)) + + nonzero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + nonzero_monitor.global_rank = 1 + nonzero_monitor.world_size = 2 + nonzero_monitor.gloo_group = sentinel_group + + def nonroot_gather(input_tensor, gather_list, dst, group): + assert input_tensor is local_round_stats + assert gather_list is None + assert dst == 0 and group is sentinel_group + + monkeypatch.setattr(monitor_module.dist, "gather", nonroot_gather) + assert nonzero_monitor._gather_round_stats(local_round_stats) is None + + +def test_find_fused_moe_weights_discovers_any_layer_member_once_and_sorts(monkeypatch): + class FakeFusedMoeWeight: + def __init__(self, layer_num, enabled=True): + self.layer_num_ = layer_num + self.enable_ep_moe = enabled + + monkeypatch.setattr(monitor_module, "FusedMoeWeight", FakeFusedMoeWeight) + first = FakeFusedMoeWeight(3) + second = FakeFusedMoeWeight(1) + disabled = FakeFusedMoeWeight(0, enabled=False) + model = SimpleNamespace( + trans_layers_weight=[ + SimpleNamespace(experts_=first, alias=first, ignored=disabled), + SimpleNamespace(any_direct_member=second), + ] + ) + + assert monitor_module._find_fused_moe_weights(model) == [second, first] + + +def test_monitor_disable_detaches_counters_from_all_impls(): + impl = SimpleNamespace(ep_balance_counters="unset") + monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) + monitor.weights = [SimpleNamespace(fuse_moe_impl=impl)] + monitor.enabled = True + monitor._disable() + assert impl.ep_balance_counters is None + assert not monitor.enabled + + +def test_critical_overhead_requires_minimum_samples_for_every_layer(): + round_stats = torch.tensor([[[[200, 32], [200, 32]], [[1, 32], [1, 32]]]], dtype=torch.int64) + assert ( + calculate_prefill_balance_stats( + round_stats, + layer_routed_experts=torch.tensor([1, 1]), + layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), + layer_topks=torch.tensor([2.0, 2.0]), + source_token_replication=1, + ) + is None + ) + + +def test_ep_moe_normal_and_prefill_enable_monitor_by_default(): + assert should_enable_ep_balance_monitor(_monitor_args()) + assert should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill")) + + +def test_disable_ep_balance_monitor_turns_monitor_off(): + assert not should_enable_ep_balance_monitor(_monitor_args(disable_ep_balance_monitor=True)) + + +def test_non_ep_moe_and_decode_mode_do_not_enable_monitor(): + assert not should_enable_ep_balance_monitor(_monitor_args(enable_ep_moe=False)) + assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="decode")) + + +def test_prefill_cudagraph_silently_disables_monitor(): + assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill", enable_prefill_cudagraph=True)) + + +def test_sm100_silently_disables_monitor(monkeypatch): + monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: True) + assert not should_enable_ep_balance_monitor(_monitor_args()) From fbe49a2227d714ce373a0fde9d5701c77368ed85 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 20 Aug 2026 06:58:49 +0000 Subject: [PATCH 117/214] add dp cache aware --- lightllm/server/api_cli.py | 4 +- lightllm/server/core/objs/start_args_type.py | 2 +- .../router/req_queue/dp_balancer/__init__.py | 5 + .../req_queue/dp_balancer/cache_aware.py | 177 ++++++++++++++++++ 4 files changed, 185 insertions(+), 3 deletions(-) create mode 100644 lightllm/server/router/req_queue/dp_balancer/cache_aware.py diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 0206bd863b..d3b32b6811 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -257,8 +257,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--dp_balancer", type=str, default="bs_balancer", - choices=["round_robin", "bs_balancer"], - help="the dp balancer type, default is bs_balancer", + choices=["round_robin", "bs_balancer", "cache_aware"], + help="the DP balancer type; cache_aware adds token-prefix affinity, default is bs_balancer", ) parser.add_argument( "--max_req_total_len", diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index ecfb111f26..f8ed52f96f 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -212,7 +212,7 @@ class StartArgs: multinode_httpmanager_port: int = field(default=12345) disable_shm_warning: bool = field(default=False) - dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer"]}) + dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer", "cache_aware"]}) enable_fused_shared_experts: bool = field(default=False) enable_mps: bool = field(default=False) multinode_router_gloo_port: int = field(default=20001) diff --git a/lightllm/server/router/req_queue/dp_balancer/__init__.py b/lightllm/server/router/req_queue/dp_balancer/__init__.py index 34f994f8a2..24db8fcd5f 100644 --- a/lightllm/server/router/req_queue/dp_balancer/__init__.py +++ b/lightllm/server/router/req_queue/dp_balancer/__init__.py @@ -2,6 +2,7 @@ from typing import List from lightllm.server.router.req_queue.base_queue import BaseQueue from .bs import DpBsBalancer +from .cache_aware import DpCacheAwareBalancer def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): @@ -9,5 +10,9 @@ def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): return RoundRobinDpBalancer(dp_size_in_node, inner_queues) elif args.dp_balancer == "bs_balancer": return DpBsBalancer(dp_size_in_node, inner_queues) + elif args.dp_balancer == "cache_aware": + if args.disable_dynamic_prompt_cache: + raise ValueError("cache_aware DP balancing requires dynamic prompt cache") + return DpCacheAwareBalancer(dp_size_in_node, inner_queues) else: raise ValueError(f"Invalid dp balancer: {args.dp_balancer}") diff --git a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py new file mode 100644 index 0000000000..bbc24c0ff6 --- /dev/null +++ b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py @@ -0,0 +1,177 @@ +"""DP-local cache-affinity routing based on bounded token-prefix history. + +The router owns this heuristic index. It records dispatch history rather than querying +the infer processes' radix trees, so stale entries can only affect placement, not KV +cache correctness. +""" + +from __future__ import annotations + +import random +from collections import OrderedDict +from dataclasses import dataclass +from typing import List, Optional, Tuple + +import xxhash + +from lightllm.server.router.batch import Batch, Req +from lightllm.server.router.req_queue.base_queue import BaseQueue + +from .base import DpBalancer + + +PrefixHash = Tuple[int, int] + + +@dataclass(slots=True) +class DpCacheAwareConfig: + cache_threshold: float = 0.5 + balance_rel_threshold: float = 1.8 + # DeepSeek-V4 prompt-cache entries are reusable only at 256-token boundaries. + block_size: int = 256 + max_cache_entries: int = 1_000_000 + evict_entries: int = 10_000 + + +class TokenPrefixCache: + """Bounded LRU mapping from cumulative token-prefix hashes to local DP indexes.""" + + def __init__(self, block_size: int, max_entries: int, evict_entries: int) -> None: + if block_size < 1: + raise ValueError(f"block_size must be >= 1, got {block_size}") + if max_entries < 0: + raise ValueError(f"max_entries must be >= 0, got {max_entries}") + if evict_entries < 1: + raise ValueError(f"evict_entries must be >= 1, got {evict_entries}") + self.block_size = block_size + self.max_entries = max_entries + self.evict_entries = evict_entries + self._prefix_to_dp: OrderedDict[int, int] = OrderedDict() + + def hash_prefixes(self, prompt_ids) -> List[PrefixHash]: + cacheable_token_count = max(0, len(prompt_ids) - 1) + cacheable_token_count = cacheable_token_count // self.block_size * self.block_size + if cacheable_token_count == 0: + return [] + + token_view = memoryview(prompt_ids) + item_size = token_view.itemsize + prompt_bytes = token_view.cast("B") + # A collision only changes a routing hint; it cannot affect cache correctness. + # xxh3-64 keeps the 1M-entry index compact and hashes 1M-token prompts faster. + hasher = xxhash.xxh3_64() + prefix_hashes = [] + for start in range(0, cacheable_token_count, self.block_size): + end = start + self.block_size + hasher.update(prompt_bytes[start * item_size : end * item_size]) + prefix_hashes.append((hasher.intdigest(), end)) + return prefix_hashes + + def match(self, prefix_hashes: List[PrefixHash]) -> Tuple[Optional[int], int]: + for prefix_hash, token_count in reversed(prefix_hashes): + try: + dp_index = self._prefix_to_dp[prefix_hash] + except KeyError: + continue + self._prefix_to_dp.move_to_end(prefix_hash) + return dp_index, token_count + return None, 0 + + def insert(self, prefix_hashes: List[PrefixHash], dp_index: int, start_index: int = 0) -> None: + for prefix_index in range(start_index, len(prefix_hashes)): + prefix_hash = prefix_hashes[prefix_index][0] + self._prefix_to_dp[prefix_hash] = dp_index + self._prefix_to_dp.move_to_end(prefix_hash) + + if len(self._prefix_to_dp) > self.max_entries: + evict_count = len(self._prefix_to_dp) - self.max_entries + self.evict_entries + for _ in range(min(evict_count, len(self._prefix_to_dp))): + self._prefix_to_dp.popitem(last=False) + + def __len__(self) -> int: + return len(self._prefix_to_dp) + + +class DpCacheAwareBalancer(DpBalancer): + """Route matching token prefixes to the same local DP unless load requires rebalancing.""" + + def __init__( + self, + dp_size_in_node: int, + inner_queues: List[BaseQueue], + config: Optional[DpCacheAwareConfig] = None, + ) -> None: + super().__init__(dp_size_in_node, inner_queues) + self.config = config or DpCacheAwareConfig() + self.prefix_cache = TokenPrefixCache( + block_size=self.config.block_size, + max_entries=self.config.max_cache_entries, + evict_entries=self.config.evict_entries, + ) + + def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: List[List[Req]]) -> None: + if not reqs_waiting_for_dp_index: + return + + current_load_per_dp = [0 for _ in range(self.dp_size_in_node)] + if current_batch is not None: + current_load_per_dp = current_batch.get_all_dp_req_num() + total_load_per_dp = [ + current_load_per_dp[dp_index] + len(self.inner_queues[dp_index].waiting_req_list) + for dp_index in range(self.dp_size_in_node) + ] + + for req_group in reqs_waiting_for_dp_index: + first_req = req_group[0] + linked_prompt_ids = False + if not hasattr(first_req, "shm_prompt_ids"): + first_req.link_prompt_ids_shm_array() + linked_prompt_ids = True + try: + prefix_hashes = self.prefix_cache.hash_prefixes(first_req.get_prompt_ids_numpy()) + finally: + if linked_prompt_ids: + first_req.shm_prompt_ids.detach_shm() + del first_req.shm_prompt_ids + + cache_dp_index = None + matched_token_count = 0 + if not first_req.sample_params.disable_prompt_cache: + matched_dp_index, matched_token_count = self.prefix_cache.match(prefix_hashes) + match_rate = matched_token_count / first_req.input_len if first_req.input_len else 0.0 + if match_rate > self.config.cache_threshold: + cache_dp_index = matched_dp_index + + idle_dp_indexes = [dp_index for dp_index, load in enumerate(total_load_per_dp) if load == 0] + if idle_dp_indexes: + if cache_dp_index in idle_dp_indexes: + selected_dp_index = cache_dp_index + else: + selected_dp_index = random.choice(idle_dp_indexes) + else: + min_load = min(total_load_per_dp) + least_loaded_dp_indexes = [ + dp_index for dp_index, load in enumerate(total_load_per_dp) if load == min_load + ] + least_loaded_dp_index = random.choice(least_loaded_dp_indexes) + if cache_dp_index is None: + selected_dp_index = least_loaded_dp_index + else: + group_load = len(req_group) + cache_projected_load = total_load_per_dp[cache_dp_index] + group_load + least_projected_load = total_load_per_dp[least_loaded_dp_index] + group_load + if cache_projected_load > least_projected_load * self.config.balance_rel_threshold: + selected_dp_index = least_loaded_dp_index + else: + selected_dp_index = cache_dp_index + + for req in req_group: + req.sample_params.suggested_dp_index = selected_dp_index + self.inner_queues[selected_dp_index].extend(req_group) + total_load_per_dp[selected_dp_index] += len(req_group) + insert_start_index = 0 + if cache_dp_index == selected_dp_index: + insert_start_index = (matched_token_count + self.config.block_size - 1) // self.config.block_size + self.prefix_cache.insert(prefix_hashes, selected_dp_index, start_index=insert_start_index) + + reqs_waiting_for_dp_index.clear() From 085b05c3429b655ab54840bb50604dec210fdf9a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 21 Aug 2026 08:40:30 +0000 Subject: [PATCH 118/214] - check full KV, SWA, C4, and C128 capacity before allocation - keep resource-constrained requests pending for retry - rematch radix cache on retry and release failed match references --- .../pd/decode_node_impl/decode_impl.py | 73 +++++++++++++++++-- 1 file changed, 65 insertions(+), 8 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 1c2828f89f..fb545eb99e 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -53,7 +53,8 @@ def _post_init_reqs(self, uninit_reqs: List[InferReq]): for req_obj in uninit_reqs: req_obj: InferReq = req_obj # for easy typing # 构建 chuncked trans task - self._decode_node_gen_trans_tasks(req_obj=req_obj) + if not self._decode_node_gen_trans_tasks(req_obj=req_obj): + PDDecodeNode._drop_pending_prompt_cache(self, req_obj) return @@ -67,6 +68,13 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: for request_id in req_ids: req_obj: InferReq = g_infer_context.requests_mapping[request_id] + # pending 期间优先重新匹配 radix;准入失败时释放引用,留待下轮重试。 + if req_obj.pd_task_num == 0 and not req_obj.infer_aborted: + req_obj._match_radix_cache() + if not self._decode_node_gen_trans_tasks(req_obj=req_obj): + PDDecodeNode._drop_pending_prompt_cache(self, req_obj) + continue + if self.is_master_in_dp and req_obj.infer_aborted and req_obj.pd_task_num != 0: self.info_queue.put(PDAbortReq(request_id=req_obj.req_id, device_id=req_obj.pd_trans_device_id)) @@ -112,9 +120,22 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: ans_list.append(req_obj) return ans_list - def _decode_node_gen_trans_tasks(self, req_obj: InferReq): + def _drop_pending_prompt_cache(self, req_obj: InferReq) -> None: + """准入失败时撤销 D 侧命中,避免 pending 请求长期占用 radix 引用。""" + assert req_obj.pd_task_num == 0 + if req_obj.shared_kv_node is not None: + self.radix_cache.dec_node_ref_counter(req_obj.shared_kv_node) + req_obj.shared_kv_node = None + # 借用的 full slots 仍由 radix 持有;下一轮从 0 传输时会覆盖请求表。 + req_obj.cur_kv_len = 0 + req_obj.pd_trans_kv_start_index = 0 + req_obj.shm_req.shm_cur_kv_len = 0 + req_obj.shm_req.prompt_cache_len = 0 + return + + def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: """ - decode node 生成所有的传输任务对象。 + decode node 生成所有的传输任务对象;资源不足时返回 False 留待下轮重试。 """ group = PDChunckedTransTaskGroup() input_len = req_obj.shm_req.input_len @@ -125,15 +146,51 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): need_mem_size = input_len - req_obj.cur_kv_len if need_mem_size > 0: - if self.radix_cache is not None: + req_manager = self.model.req_manager + is_dsv4_req_manager = isinstance(req_manager, DeepseekV4ReqManager) + if is_dsv4_req_manager: + mem_manager = req_manager.mem_manager + ready_len = req_obj.cur_kv_len + + if need_mem_size > g_infer_context.get_can_alloc_token_num(): + return False + if self.radix_cache is not None: + self.radix_cache.free_radix_cache_to_get_enough_token(need_mem_size) + if need_mem_size > mem_manager.allocator.can_use_mem_size: + return False + + prompt_page = req_manager.get_prompt_cache_page_size() + swa_start = max(ready_len, max(0, input_len // prompt_page * prompt_page - prompt_page)) + swa_page = mem_manager.swa_pool.page_size + swa_need = max(0, (input_len - 1) // swa_page - (swa_start + swa_page - 1) // swa_page + 1) + + _, c4_need, c128_need = req_obj.get_dsv4_prefill_need_page_and_slot_num(is_chuncked_prefill=False) + c4_allocator = mem_manager.c4_page_allocator + c128_allocator = mem_manager.c128_allocator + + # D ingress 会立即分配派生槽,必须在任何分配前先兑现并检查实际容量。 + if self.radix_cache is not None: + self.radix_cache.free_radix_cache_to_get_enough_c4_pages(c4_need) + self.radix_cache.free_radix_cache_to_get_enough_c128_slots(c128_need) + swa_shortage = swa_need - mem_manager.swa_page_allocator.can_use_mem_size + if swa_shortage > 0: + self.radix_cache.free_unreferenced_swa_pages(swa_shortage) + + if ( + swa_need > mem_manager.swa_page_allocator.can_use_mem_size + or (c4_allocator is not None and c4_need > c4_allocator.can_use_mem_size) + or (c128_allocator is not None and c128_need > c128_allocator.can_use_mem_size) + ): + return False + + if self.radix_cache is not None and not is_dsv4_req_manager: self.radix_cache.free_radix_cache_to_get_enough_token(need_mem_size) - req_manager = self.model.req_manager mem_indexes = req_manager.mem_manager.alloc(need_size=need_mem_size) req_manager.req_to_token_indexs[ req_obj.req_idx, req_obj.cur_kv_len : (req_obj.cur_kv_len + need_mem_size) ] = mem_indexes - if isinstance(req_manager, DeepseekV4ReqManager): + if is_dsv4_req_manager: req_manager.prepare_pd_decode_cache( req_list=[req_obj.req_idx], ready_list=[req_obj.cur_kv_len], @@ -169,7 +226,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): group=group, page_kind="linear_att_state", ) - if isinstance(req_manager, DeepseekV4ReqManager): + if is_dsv4_req_manager: torch.cuda.current_stream().synchronize() else: assert req_obj.cur_kv_len == input_len - 1 @@ -187,7 +244,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq): if self.is_master_in_dp: self.info_queue.put(group) - return + return True def _create_pd_trans_task( self, From 7c65cd00ac66f0998f64709945de18c7139bb3d2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 21 Aug 2026 08:43:48 +0000 Subject: [PATCH 119/214] fix(pd): preserve aborts received before request registration --- lightllm/server/httpserver/manager.py | 23 ++++++++++++++++++++++- lightllm/server/httpserver/pd_loop.py | 13 ++++--------- 2 files changed, 26 insertions(+), 10 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index f92ee84ec2..89fd01bcfc 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -110,6 +110,8 @@ def __init__( self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) self.req_id_to_out_inf: Dict[int, ReqStatus] = {} # value type (out_str, metadata, finished, event) + # key 存在表示 PD 请求正在登记,value 表示登记期间是否已收到 ABORT。 + self._pd_registration_abort_flags: Dict[int, bool] = {} self.forwarding_queue: AsyncQueue = None # p d 分离模式使用的转发队列, 需要延迟初始化 self.max_req_total_len = args.max_req_total_len @@ -312,6 +314,21 @@ def alloc_req_id(self, sampling_params): assert False, "dead code path" return group_request_id + def begin_pd_request_registration(self, group_req_id: int) -> None: + self._pd_registration_abort_flags[group_req_id] = False + + def cancel_pd_request_registration(self, group_req_id: int) -> None: + """PD 请求在正式登记前结束时,清理对应的待处理 ABORT。""" + self._pd_registration_abort_flags.pop(group_req_id, None) + + def _register_req_status(self, group_req_id: int, req_status: "ReqStatus") -> None: + """发布 PD 请求,并立即消费登记期间收到的 ABORT。""" + self.req_id_to_out_inf[group_req_id] = req_status + if self._pd_registration_abort_flags.pop(group_req_id, False): + for req in req_status.group_req_objs.shm_req_objs: + req.is_aborted = True + logger.warning(f"applied pending abort for group_request_id {group_req_id}") + async def generate( self, prompt: Union[str, List[int]], @@ -450,7 +467,7 @@ async def generate( ) req_status = ReqStatus(group_request_id, multimodal_params, req_objs, start_time) - self.req_id_to_out_inf[group_request_id] = req_status + self._register_req_status(group_request_id, req_status) # RL:请求已登记到 req_id_to_out_inf 并即将转发下游,从 admission gate # 注销,避免 pause 统计里仍把它算作“等待准入”的 pending 请求。 if self.rl_controller is not None: @@ -825,6 +842,10 @@ async def _wait_to_token_package( async def abort(self, group_req_id: int) -> bool: req_status: ReqStatus = self.req_id_to_out_inf.get(group_req_id, None) if req_status is None: + if group_req_id in self._pd_registration_abort_flags: + self._pd_registration_abort_flags[group_req_id] = True + logger.warning(f"deferred abort for registering group_request_id {group_req_id}") + return True logger.warning(f"aborted group_request_id {group_req_id} not exist") return False diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index d55631e3ff..a4be1476d0 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -129,6 +129,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O group_req_id = sampling_params.group_request_id pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event + manager.begin_pd_request_registration(group_req_id) asyncio.create_task( _pd_process_generate( manager=manager, @@ -143,15 +144,7 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O elif obj[0] == ObjType.ABORT: group_req_id = obj[1] logger.warning(f"recv cmd aborted req id {group_req_id}") - if not (await manager.abort(group_req_id)): - - async def delayed_abort_task(group_req_id, retry_count): - for _ in range(retry_count): - await asyncio.sleep(5.0) - if await manager.abort(group_req_id): - break - - asyncio.create_task(delayed_abort_task(group_req_id=group_req_id, retry_count=4)) + await manager.abort(group_req_id) elif obj[0] == ObjType.PD_REQ_DECODE_NODE_INFO: _, group_req_id, decode_node_info = obj @@ -253,6 +246,8 @@ async def _pd_process_generate( logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") except BaseException as e: logger.error(str(e)) + finally: + manager.cancel_pd_request_registration(sampling_params.group_request_id) # 转发token的task From 34a5a6d136ba8c0eebc33259995d930eca930353 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 21 Aug 2026 09:23:54 +0000 Subject: [PATCH 120/214] fix(dsv4): only protect the last SWA page of active radix matches --- .../router/dynamic_prompt/radix_cache.py | 37 ++++++++++++++++--- 1 file changed, 31 insertions(+), 6 deletions(-) diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index 4e82dbf1d4..cd3d7f9773 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -174,6 +174,15 @@ def _node_swa_pages_num(self, node: TreeNode) -> int: return 0 return int(valid.sum().item()) * self._swa_pages_per_prompt_page + def _node_direct_ref_num(self, node: TreeNode) -> int: + # 子树引用同时计入父子节点,差值就是直接命中本节点的请求数。 + return node.ref_counter - sum(child.ref_counter for child in node.children.values()) + + def _node_last_valid_swa_pages_num(self, node: TreeNode) -> int: + if node.token_extra_value is None: + return 0 + return self._swa_pages_per_prompt_page if node.token_extra_value.swa_last_valid_page >= 0 else 0 + def _align_len(self, length: int) -> int: if self.page_size <= 1: return int(length) @@ -352,6 +361,8 @@ def match_prefix(self, key, update_refs=False): ans_value_list = [] tree_node = self._match_prefix_helper(self.root_node, key, ans_value_list, update_refs=update_refs) if tree_node != self.root_node: + if update_refs and self._swa_pages_per_prompt_page > 0 and self._node_direct_ref_num(tree_node) == 1: + self.swa_refed_pages_num += self._node_last_valid_swa_pages_num(tree_node) if len(ans_value_list) != 0: value = torch.concat(ans_value_list) else: @@ -427,7 +438,6 @@ def _match_prefix_helper_no_recursion( # from 0 to 1 need update refs token num if node.ref_counter == 1: self.refed_tokens_num.arr[0] += len(node.token_mem_index_value) - self.swa_refed_pages_num += self._node_swa_pages_num(node) if len(key) == 0: return node @@ -457,7 +467,6 @@ def _match_prefix_helper_no_recursion( # from 0 to 1 need update refs token num if split_parent_node.ref_counter == 1: self.refed_tokens_num.arr[0] += len(split_parent_node.token_mem_index_value) - self.swa_refed_pages_num += self._node_swa_pages_num(split_parent_node) if child.is_leaf(): self.evict_tree_set.add(child) @@ -581,13 +590,18 @@ def dec_node_ref_counter(self, node: TreeNode): return # 如果减引用的是叶节点,需要先从 evict_tree_set 中移除 old_node = node + if ( + self._swa_pages_per_prompt_page > 0 + and old_node is not self.root_node + and self._node_direct_ref_num(old_node) == 1 + ): + self.swa_refed_pages_num -= self._node_last_valid_swa_pages_num(old_node) if old_node.is_leaf(): self.evict_tree_set.discard(old_node) while node is not None: if node.ref_counter == 1: self.refed_tokens_num.arr[0] -= len(node.token_mem_index_value) - self.swa_refed_pages_num -= self._node_swa_pages_num(node) node.ref_counter -= 1 node = node.parent @@ -601,13 +615,18 @@ def add_node_ref_counter(self, node: TreeNode): return # 如果减引用的是叶节点,需要先从 evict_tree_set 中移除 old_node = node + if ( + self._swa_pages_per_prompt_page > 0 + and old_node is not self.root_node + and self._node_direct_ref_num(old_node) == 0 + ): + self.swa_refed_pages_num += self._node_last_valid_swa_pages_num(old_node) if old_node.is_leaf(): self.evict_tree_set.discard(old_node) while node is not None: if node.ref_counter == 0: self.refed_tokens_num.arr[0] += len(node.token_mem_index_value) - self.swa_refed_pages_num += self._node_swa_pages_num(node) node.ref_counter += 1 node = node.parent @@ -665,7 +684,7 @@ def _print_helper(self, node: TreeNode, indent): return def free_unreferenced_swa_pages(self, need_pages: int) -> None: - """DeepSeek-V4 swa free hook: 页 allocator 触底时,回收 ref_count==0 节点的 swa 页。""" + """DeepSeek-V4 swa free hook: 页 allocator 触底时回收未被命中终点保护的 swa 页。""" if self.mem_manager is None or self.extra_value_ops is None: return allocator = self.mem_manager.swa_page_allocator @@ -679,13 +698,15 @@ def free_unreferenced_swa_pages(self, need_pages: int) -> None: if allocator.can_use_mem_size + evict_swa_pages >= target: break node = leaf - while node is not None and node is not self.root_node and node.ref_counter == 0: + while node is not None and node is not self.root_node: node_id = id(node) if node_id in visited: node = node.parent continue visited.add(node_id) + has_direct_ref = self._node_direct_ref_num(node) > 0 + payload = node.token_extra_value if ( len(node.token_mem_index_value) > 0 @@ -695,6 +716,10 @@ def free_unreferenced_swa_pages(self, need_pages: int) -> None: last_page = int(payload.swa_last_valid_page) if last_page >= 0: if free_last: + if has_direct_ref: + # 活跃请求恰好命中该节点时,最后一页仍需保留。 + node = node.parent + continue page_slice = slice(last_page, last_page + 1) else: page_slice = slice(0, last_page) From aca684cfe8bbbfc168d704f6eb6cdfd561fe043e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 21 Aug 2026 09:53:11 +0000 Subject: [PATCH 121/214] support PreTrainedTokenizerFast --- lightllm/server/tokenizer.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/lightllm/server/tokenizer.py b/lightllm/server/tokenizer.py index 6aea7cd672..1e7522ad66 100644 --- a/lightllm/server/tokenizer.py +++ b/lightllm/server/tokenizer.py @@ -65,6 +65,11 @@ def get_tokenizer( try: tokenizer = AutoTokenizer.from_pretrained(tokenizer_name, trust_remote_code=trust_remote_code, *args, **kwargs) + except ValueError as e: + if tokenizer_mode == "slow" or "Tokenizer class TokenizersBackend does not exist" not in str(e): + raise + logger.warning("Transformers does not provide TokenizersBackend; loading tokenizer.json as a fast tokenizer") + tokenizer = PreTrainedTokenizerFast.from_pretrained(tokenizer_name, *args, **kwargs) except TypeError as e: # The LLaMA tokenizer causes a protobuf error in some environments, using slow mode. # you can try pip install protobuf==3.20.0 to try repair From 7445e42d7662477f947fd327db77403b6e57b41c Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Fri, 21 Aug 2026 19:36:19 +0800 Subject: [PATCH 122/214] fix: fix DeepEP workspace error --- .../fused_moe/grouped_fused_moe_ep.py | 6 +- lightllm/distributed/communication_op.py | 199 +++++++++++------- 2 files changed, 129 insertions(+), 76 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 210da61da1..1c6feea31b 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -553,8 +553,10 @@ def chunked_expanded_moe_forward( if max_chunk_rows == 0: raise RuntimeError( - f"DeepEP workspace with {workspace.numel()} bytes cannot hold the dense output and " - f"one {alignment}-row temporary chunk" + "RDMA workspace sizing invariant violated: " + f"workspace_bytes={workspace.numel()}, gather_rows={gather_rows}, " + f"hidden_size={hidden_size}, intermediate_size={intermediate_size}, " + f"chunk_rows={alignment}" ) max_chunk_rows = min(all_tokens, max_chunk_rows) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 330db23bb1..f4d57bdd1b 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -18,9 +18,11 @@ # limitations under the License. +import math import os import torch import torch.distributed as dist +from dataclasses import dataclass from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union from lightllm.utils.log_utils import init_logger @@ -43,54 +45,94 @@ logger = init_logger(__name__) -def get_deep_ep_prefill_moe_workspace_size( - num_max_tokens_per_rank: int, +_TENSOR_BUFFER_ALIGNMENT_BYTES = 256 +_DEEPEP_PREFILL_CHUNK_ROWS = 128 +_DEEPEP_GATHER_ROWS_CACHE_ALIGNMENT = 1024 + + +@dataclass(frozen=True) +class _LegacyLowLatencyRdmaSizing: + """Capacity required to reuse each legacy DeepEP RDMA slice in prefill.""" + + num_rdma_bytes: int + per_workspace_bytes: int + prefill_required_total_bytes: int + max_dense_rows: int + microbatch_count: int + + +def _align_up(value: int, alignment: int) -> int: + return (value + alignment - 1) // alignment * alignment + + +def _calculate_legacy_low_latency_rdma_sizing( + decode_size_hint: int, + global_world_size: int, + prefill_tokens_per_rank: int, hidden_size: int, - intermediate_size: int, - num_experts_per_tok: int, - num_experts: int, - world_size: int, + moe_intermediate_size: Optional[int], hidden_dtype: torch.dtype, -) -> int: - """Size one Prefill microbatch workspace for the bounded grouped-MoE path. - - The target chunk covers balanced expert routing in one pass. More skewed - routing remains correct because the consumer already splits expanded rows. + microbatch_count: int, +) -> _LegacyLowLatencyRdmaSizing: + """Return RDMA storage large enough for one prefill chunk in every slice. + + The grouped FP8 expanded-MoE path keeps the dense gather output alive, + then executes a 128-row chunk through W1, quantization, and W2. The byte + accounting mirrors TensorBufferManager's 256-byte first-fit allocation and + the actual tensor lifetimes in ``chunked_expanded_moe_forward``. """ - tensor_alignment = 256 - expert_alignment = 128 - metadata_row_granularity = 1024 - - assert num_experts % world_size == 0 - assert intermediate_size % expert_alignment == 0 - - def align_up(value: int, alignment: int) -> int: - return (value + alignment - 1) // alignment * alignment - - hidden_bytes = torch.empty((), dtype=hidden_dtype).element_size() - max_gather_rows = align_up(world_size * num_max_tokens_per_rank, metadata_row_granularity) - num_local_experts = num_experts // world_size - target_chunk_rows = align_up( - num_max_tokens_per_rank * num_experts_per_tok + num_local_experts * (expert_alignment - 1), - expert_alignment, - ) - - gather_out = align_up(max_gather_rows * hidden_size * hidden_bytes, tensor_alignment) - silu_out = align_up(target_chunk_rows * intermediate_size * hidden_bytes, tensor_alignment) - gemm_out_a = align_up(target_chunk_rows * 2 * intermediate_size * hidden_bytes, tensor_alignment) - quant_out = align_up(target_chunk_rows * intermediate_size, tensor_alignment) - quant_scale = align_up(target_chunk_rows * (intermediate_size // expert_alignment) * 4, tensor_alignment) - gemm_out_b = align_up(target_chunk_rows * hidden_size * hidden_bytes, tensor_alignment) + if moe_intermediate_size is None: + raise ValueError("Legacy DeepEP low-latency buffer requires moe_intermediate_size") + if global_world_size <= 0 or prefill_tokens_per_rank <= 0 or hidden_size <= 0: + raise ValueError("DeepEP workspace dimensions must be positive") + if moe_intermediate_size <= 0 or moe_intermediate_size % _DEEPEP_PREFILL_CHUNK_ROWS: + raise ValueError("moe_intermediate_size must be a positive multiple of 128 for legacy DeepEP FP8 MoE") + if microbatch_count <= 0: + raise ValueError("DeepEP microbatch_count must be positive") + + def tensor_bytes(*shape: int, itemsize: int) -> int: + return _align_up(math.prod(shape) * itemsize, _TENSOR_BUFFER_ALIGNMENT_BYTES) + + chunk_rows = _DEEPEP_PREFILL_CHUNK_ROWS + hidden_itemsize = hidden_dtype.itemsize + block_size_k = 128 + scale_cols = moe_intermediate_size // block_size_k + silu_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, itemsize=hidden_itemsize) + gemm_out_a_bytes = tensor_bytes(chunk_rows, 2 * moe_intermediate_size, itemsize=hidden_itemsize) + # Legacy expanded MoE quantizes activations to FP8, i.e. one byte per value. + quant_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, itemsize=1) + # HAS_SGL_KERNEL uses [scale_cols, chunk_rows], the fallback [chunk_rows, + # scale_cols]; their element counts are identical. + scale_bytes = tensor_bytes(chunk_rows, scale_cols, itemsize=torch.float32.itemsize) + gemm_out_b_bytes = tensor_bytes(chunk_rows, hidden_size, itemsize=hidden_itemsize) + + w1_peak_bytes = silu_bytes + gemm_out_a_bytes + quant_peak_bytes = silu_bytes + quant_bytes + scale_bytes + # A is freed before Q/L. S is freed before B; first-fit can reuse S only + # when B fits there, otherwise B must fit as a contiguous tail allocation. + if silu_bytes >= gemm_out_b_bytes: + temp_bytes = max(w1_peak_bytes, quant_peak_bytes) + else: + temp_bytes = max(w1_peak_bytes, quant_peak_bytes + gemm_out_b_bytes) - # TensorBufferManager uses first-fit allocation. W1 keeps silu_out and gemm_out_a - # together; after gemm_out_a is reused by quantization, the freed silu_out block - # can only hold gemm_out_b when it is large enough. - w1_peak = gather_out + silu_out + gemm_out_a - w2_peak = gather_out + silu_out + quant_out + quant_scale - if gemm_out_b > silu_out: - w2_peak += gemm_out_b - # TensorBufferManager may trim an unaligned prefix from the supplied view. - return max(w1_peak, w2_peak) + tensor_alignment + max_dense_rows = _align_up( + global_world_size * prefill_tokens_per_rank, + _DEEPEP_GATHER_ROWS_CACHE_ALIGNMENT, + ) + gather_bytes = tensor_bytes(max_dense_rows, hidden_size, itemsize=hidden_itemsize) + per_workspace_bytes = gather_bytes + temp_bytes + prefill_required_total_bytes = per_workspace_bytes * microbatch_count + num_rdma_bytes = _align_up( + max(decode_size_hint, prefill_required_total_bytes), + microbatch_count * _TENSOR_BUFFER_ALIGNMENT_BYTES, + ) + return _LegacyLowLatencyRdmaSizing( + num_rdma_bytes=num_rdma_bytes, + per_workspace_bytes=per_workspace_bytes, + prefill_required_total_bytes=prefill_required_total_bytes, + max_dense_rows=max_dense_rows, + microbatch_count=microbatch_count, + ) try: @@ -294,14 +336,38 @@ def new_deepep_group( enable_mega_moe_buffer = False has_legacy_moe_layer = True - enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" - enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" + # The legacy buffer is also the only prefill workspace. Pure prefill + # does not use its low-latency protocol, but still needs its RDMA + # storage for the bounded grouped-MoE compute path. + enable_low_latency_buffer = has_legacy_moe_layer if enable_low_latency_buffer: - # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 - # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 - num_rdma_bytes = deep_ep.Buffer.get_low_latency_rdma_size_hint( - self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts + if self.ll_decode_num_tokens is None: + decode_size_hint = 0 + else: + decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( + self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts + ) + rdma_sizing = _calculate_legacy_low_latency_rdma_sizing( + decode_size_hint=decode_size_hint, + global_world_size=global_world_size, + prefill_tokens_per_rank=prefill_num_max_dispatch_tokens_per_rank, + hidden_size=hidden_size, + moe_intermediate_size=moe_intermediate_size, + hidden_dtype=get_torch_dtype(args.data_type), + microbatch_count=len(self.groups), + ) + num_rdma_bytes = rdma_sizing.num_rdma_bytes + logger.info( + "Initialize DeepEP legacy RDMA workspace: decode_size_hint=%s, " + "prefill_required_total_bytes=%s, num_rdma_bytes=%s, " + "per_workspace_bytes=%s, max_dense_rows=%s, microbatch_count=%s", + decode_size_hint, + rdma_sizing.prefill_required_total_bytes, + num_rdma_bytes, + rdma_sizing.per_workspace_bytes, + rdma_sizing.max_dense_rows, + rdma_sizing.microbatch_count, ) self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, @@ -313,22 +379,6 @@ def new_deepep_group( torch.uint8, use_rdma_buffer=True ) - if enable_prefill_workspace: - workspace_size = get_deep_ep_prefill_moe_workspace_size( - num_max_tokens_per_rank=self.ll_num_tokens, - hidden_size=self.ll_hidden, - intermediate_size=moe_intermediate_size, - num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_num_experts, - world_size=global_world_size, - hidden_dtype=get_torch_dtype(args.data_type), - ) - self.ep_prefill_moe_workspace = torch.empty( - workspace_size * len(self.groups), - dtype=torch.uint8, - device=torch.device("cuda", torch.cuda.current_device()), - ) - if enable_mega_moe_buffer: # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy # low-latency buffer,因此纯 FP4 模型无需承担后者的大块 RDMA 显存。 @@ -375,12 +425,12 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: """Return a slice of the workspace used by DeepEP Prefill MoE kernels. - Pure Prefill nodes own a dedicated workspace and do not initialize a - low-latency Decode buffer. Decode-capable nodes reuse the idle local - RDMA storage. With one communication group, the default + All legacy FP8 MoE modes, including pure prefill, use local DeepEP + low-latency RDMA storage. Pure prefill only treats it as workspace; + decode-capable modes reuse the same storage after decode has released + its low-latency layout. With one communication group, the default ``microbatch_index=0`` receives the whole workspace. With multiple - groups, the workspace is split into ``len(self.groups)`` equal slices - and each in-flight microbatch uses the slice matching its group index. + groups, the workspace is split into ``len(self.groups)`` equal slices. Args: microbatch_index: Zero-based microbatch and communication-group @@ -404,10 +454,11 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. def clear_deepep_buffer(self): """ Decode-capable modes reuse the low-latency RDMA buffer during Prefill, - so clean it before the next low-latency Decode. Pure Prefill owns a - dedicated workspace and has no low-latency buffer to clean. + so clean it before the next low-latency Decode. Pure Prefill uses the + same RDMA storage only as workspace and has no low-latency layout to + clean. """ - if self.ep_low_latency_buffer is not None: + if self.ep_low_latency_buffer is not None and self.ll_decode_num_tokens is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts ) From 4ba4307d93903a0a8e7358aa59c2e5671723d55c Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Mon, 24 Aug 2026 07:29:07 +0800 Subject: [PATCH 123/214] feat: support disk cache --- .../mode_backend/dsv4_multi_level_kv_cache.py | 51 +++++++++++++++++-- 1 file changed, 47 insertions(+), 4 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py index 2f058af620..eef20b9a57 100644 --- a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py @@ -15,6 +15,26 @@ from .multi_level_kv_cache import MultiLevelKvCacheModule +def _split_dsv4_loaded_cache_lengths( + original_gpu_kv_len: int, + loaded_end: int, + requested_end: int, + disk_prompt_cache_len: int, +) -> tuple[int, int]: + """Split an actual CPU-cache load into CPU and disk matched token counts.""" + load_start = max(0, int(original_gpu_kv_len)) + load_end = max(load_start, int(loaded_end)) + actual_loaded_len = load_end - load_start + + matched_end = max(0, int(requested_end)) + matched_disk_len = min(max(0, int(disk_prompt_cache_len)), matched_end) + disk_start = matched_end - matched_disk_len + actual_disk_len = max(0, min(load_end, matched_end) - max(load_start, disk_start)) + actual_disk_len = min(actual_disk_len, actual_loaded_len) + actual_cpu_len = actual_loaded_len - actual_disk_len + return actual_cpu_len, actual_disk_len + + class Dsv4MultiLevelKvCacheModule(MultiLevelKvCacheModule): def __init__(self, backend): super().__init__(backend) @@ -32,9 +52,23 @@ def _try_release_dsv4_session(self, session: "Dsv4CpuStoreSession") -> None: if session.leased_pages: self.cpu_cache_client.lock.acquire_sleep1ms() try: - # Cumulative hashes make the root page the most valuable entry. - # Releasing tail-to-root makes the tail oldest in the LRU. - self.cpu_cache_client.deref_pages(list(reversed(session.leased_pages))) + if self.args.enable_disk_cache: + # A disk-cache group must contain one complete request prefix in + # root-to-tail order. Incremental store batches may complete in + # a different order, so do not publish or release any page until + # every page leased by this session is ready. + if not self.cpu_cache_client.check_allpages_ready(session.leased_pages): + return + self.cpu_cache_client.update_pages_status_to_ready( + page_list=session.leased_pages, + deref=True, + disk_offload_enable=True, + token_num_in_page_list=(len(session.leased_pages) * self.args.cpu_cache_token_page_size), + ) + else: + # Cumulative hashes make the root page the most valuable entry. + # Releasing tail-to-root makes the tail oldest in the LRU. + self.cpu_cache_client.deref_pages(list(reversed(session.leased_pages))) finally: self.cpu_cache_client.lock.release() del self._dsv4_store_sessions[session.request_id] @@ -218,6 +252,8 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): assert len(page_list) <= len(page_len_list) gpu_kv_len = int(req.cur_kv_len) + requested_end = gpu_kv_len + matched_disk_len = int(req.shm_req.disk_prompt_cache_len) if is_master_in_dp: session = Dsv4CpuStoreSession( request_id=req.req_id, @@ -291,7 +327,14 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): idle_token_num -= token_num if is_master_in_dp: - req.shm_req.cpu_prompt_cache_len = loaded_end - gpu_kv_len + cpu_prompt_cache_len, disk_prompt_cache_len = _split_dsv4_loaded_cache_lengths( + original_gpu_kv_len=gpu_kv_len, + loaded_end=loaded_end, + requested_end=requested_end, + disk_prompt_cache_len=matched_disk_len, + ) + req.shm_req.cpu_prompt_cache_len = cpu_prompt_cache_len + req.shm_req.disk_prompt_cache_len = disk_prompt_cache_len req.shm_req.shm_cur_kv_len = loaded_end session.load_submitted = True if loaded_end > gpu_kv_len: From bab7bc68ab0d600c13a1a44de39ef8558053d7c0 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 24 Aug 2026 01:43:32 +0000 Subject: [PATCH 124/214] fix(dsv4): preserve radix-owned SWA prefix when skip prefill --- lightllm/common/req_manager.py | 5 +- .../router/dynamic_prompt/radix_cache.py | 123 ++++++++++-------- 2 files changed, 74 insertions(+), 54 deletions(-) diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index a0dcdbcb8e..de68d35f49 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -593,8 +593,9 @@ def prepare_decode_swa( continue mark = self._swa_evict_marks[req_idx] if mark < 0: - # 未经过 prefill prep 的保守路径: 不回收旧位置,仅推进水位线。 - self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(seq_len - retain) + # direct-decode 中 [0, seq_len-1) 是已有 KV;exact hit 时这段前缀归 radix 所有, + # 水位必须从前缀末端开始,不能由请求回收。 + self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(seq_len - 1) continue evict_end = self._align_swa_evict_frontier(seq_len - retain) if evict_end > mark: diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index cd3d7f9773..ffc016d9e5 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -689,61 +689,80 @@ def free_unreferenced_swa_pages(self, need_pages: int) -> None: return allocator = self.mem_manager.swa_page_allocator target = allocator.can_use_mem_size + int(need_pages) - evict_slots = [] - invalidate_payloads = [] - evict_swa_pages = 0 - for free_last in (False, True): - visited = set() - for leaf in self.evict_tree_set: + planned_reclaim_pages = 0 + actual_reclaim_pages = 0 + while allocator.can_use_mem_size < target: + evict_slots = [] + invalidate_payloads = [] + evict_swa_pages = 0 + for free_last in (False, True): + visited = set() + for leaf in self.evict_tree_set: + if allocator.can_use_mem_size + evict_swa_pages >= target: + break + node = leaf + while node is not None and node is not self.root_node: + node_id = id(node) + if node_id in visited: + node = node.parent + continue + visited.add(node_id) + + has_direct_ref = self._node_direct_ref_num(node) > 0 + + payload = node.token_extra_value + if ( + len(node.token_mem_index_value) > 0 + and payload is not None + and payload.swa_page_valid is not None + ): + last_page = int(payload.swa_last_valid_page) + if last_page >= 0: + if free_last: + if has_direct_ref: + # 活跃请求恰好命中该节点时,最后一页仍需保留。 + node = node.parent + continue + page_slice = slice(last_page, last_page + 1) + else: + page_slice = slice(0, last_page) + valid_pages = int(payload.swa_page_valid[page_slice].sum().item()) + if valid_pages > 0: + start = page_slice.start * self.page_size + end = min(page_slice.stop * self.page_size, len(node.token_mem_index_value)) + if end > start: + evict_slots.append(node.token_mem_index_value[start:end]) + invalidate_payloads.append((payload, page_slice, free_last)) + evict_swa_pages += valid_pages * self._swa_pages_per_prompt_page + if allocator.can_use_mem_size + evict_swa_pages >= target: + break + node = node.parent if allocator.can_use_mem_size + evict_swa_pages >= target: break - node = leaf - while node is not None and node is not self.root_node: - node_id = id(node) - if node_id in visited: - node = node.parent - continue - visited.add(node_id) - - has_direct_ref = self._node_direct_ref_num(node) > 0 - - payload = node.token_extra_value - if ( - len(node.token_mem_index_value) > 0 - and payload is not None - and payload.swa_page_valid is not None - ): - last_page = int(payload.swa_last_valid_page) - if last_page >= 0: - if free_last: - if has_direct_ref: - # 活跃请求恰好命中该节点时,最后一页仍需保留。 - node = node.parent - continue - page_slice = slice(last_page, last_page + 1) - else: - page_slice = slice(0, last_page) - valid_pages = int(payload.swa_page_valid[page_slice].sum().item()) - if valid_pages > 0: - start = page_slice.start * self.page_size - end = min(page_slice.stop * self.page_size, len(node.token_mem_index_value)) - if end > start: - evict_slots.append(node.token_mem_index_value[start:end]) - invalidate_payloads.append((payload, page_slice, free_last)) - evict_swa_pages += valid_pages * self._swa_pages_per_prompt_page - if allocator.can_use_mem_size + evict_swa_pages >= target: - break - node = node.parent - if allocator.can_use_mem_size + evict_swa_pages >= target: + if len(evict_slots) == 0: break - if len(evict_slots) == 0: - return - self.mem_manager.evict_swa(torch.cat(evict_slots)) - for payload, page_slice, free_last in invalidate_payloads: - payload.swa_page_valid[page_slice] = False - if free_last: - payload.swa_last_valid_page = -1 - self.swa_tree_total_pages_num -= evict_swa_pages + free_pages_before = allocator.can_use_mem_size + self.mem_manager.evict_swa(torch.cat(evict_slots)) + for payload, page_slice, free_last in invalidate_payloads: + payload.swa_page_valid[page_slice] = False + if free_last: + payload.swa_last_valid_page = -1 + self.swa_tree_total_pages_num -= evict_swa_pages + planned_reclaim_pages += evict_swa_pages + actual_reclaim_pages += allocator.can_use_mem_size - free_pages_before + + # bitmap 是候选页账本,最终能否继续分配必须以物理 allocator 的真实水位为准。 + if allocator.can_use_mem_size < target or actual_reclaim_pages < planned_reclaim_pages: + logger.warning( + "DSV4 SWA reclaim mismatch: target_free_pages=%d actual_free_pages=%d " + "planned_reclaim_pages=%d actual_reclaim_pages=%d tree_pages=%d protected_pages=%d", + target, + allocator.can_use_mem_size, + planned_reclaim_pages, + actual_reclaim_pages, + self.swa_tree_total_pages_num, + self.swa_refed_pages_num, + ) return def free_radix_cache_to_get_enough_token(self, need_token_num): From 216ec51d05a27bb5ad558fb69895fc9c56d4e357 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 24 Aug 2026 02:15:08 +0000 Subject: [PATCH 125/214] revert 7445e and just use max(num_rdma_bytes, workspace_size * len(self.groups)) --- lightllm/distributed/communication_op.py | 211 +++++++++-------------- 1 file changed, 86 insertions(+), 125 deletions(-) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index f4d57bdd1b..c554213e6a 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -18,11 +18,9 @@ # limitations under the License. -import math import os import torch import torch.distributed as dist -from dataclasses import dataclass from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union from lightllm.utils.log_utils import init_logger @@ -45,95 +43,55 @@ logger = init_logger(__name__) -_TENSOR_BUFFER_ALIGNMENT_BYTES = 256 -_DEEPEP_PREFILL_CHUNK_ROWS = 128 -_DEEPEP_GATHER_ROWS_CACHE_ALIGNMENT = 1024 - - -@dataclass(frozen=True) -class _LegacyLowLatencyRdmaSizing: - """Capacity required to reuse each legacy DeepEP RDMA slice in prefill.""" - - num_rdma_bytes: int - per_workspace_bytes: int - prefill_required_total_bytes: int - max_dense_rows: int - microbatch_count: int - - -def _align_up(value: int, alignment: int) -> int: - return (value + alignment - 1) // alignment * alignment - - -def _calculate_legacy_low_latency_rdma_sizing( - decode_size_hint: int, - global_world_size: int, - prefill_tokens_per_rank: int, +def get_deep_ep_prefill_moe_workspace_size( + num_max_tokens_per_rank: int, hidden_size: int, - moe_intermediate_size: Optional[int], + intermediate_size: int, + num_experts_per_tok: int, + num_experts: int, + world_size: int, hidden_dtype: torch.dtype, - microbatch_count: int, -) -> _LegacyLowLatencyRdmaSizing: - """Return RDMA storage large enough for one prefill chunk in every slice. - - The grouped FP8 expanded-MoE path keeps the dense gather output alive, - then executes a 128-row chunk through W1, quantization, and W2. The byte - accounting mirrors TensorBufferManager's 256-byte first-fit allocation and - the actual tensor lifetimes in ``chunked_expanded_moe_forward``. - """ - if moe_intermediate_size is None: - raise ValueError("Legacy DeepEP low-latency buffer requires moe_intermediate_size") - if global_world_size <= 0 or prefill_tokens_per_rank <= 0 or hidden_size <= 0: - raise ValueError("DeepEP workspace dimensions must be positive") - if moe_intermediate_size <= 0 or moe_intermediate_size % _DEEPEP_PREFILL_CHUNK_ROWS: - raise ValueError("moe_intermediate_size must be a positive multiple of 128 for legacy DeepEP FP8 MoE") - if microbatch_count <= 0: - raise ValueError("DeepEP microbatch_count must be positive") - - def tensor_bytes(*shape: int, itemsize: int) -> int: - return _align_up(math.prod(shape) * itemsize, _TENSOR_BUFFER_ALIGNMENT_BYTES) - - chunk_rows = _DEEPEP_PREFILL_CHUNK_ROWS - hidden_itemsize = hidden_dtype.itemsize - block_size_k = 128 - scale_cols = moe_intermediate_size // block_size_k - silu_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, itemsize=hidden_itemsize) - gemm_out_a_bytes = tensor_bytes(chunk_rows, 2 * moe_intermediate_size, itemsize=hidden_itemsize) - # Legacy expanded MoE quantizes activations to FP8, i.e. one byte per value. - quant_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, itemsize=1) - # HAS_SGL_KERNEL uses [scale_cols, chunk_rows], the fallback [chunk_rows, - # scale_cols]; their element counts are identical. - scale_bytes = tensor_bytes(chunk_rows, scale_cols, itemsize=torch.float32.itemsize) - gemm_out_b_bytes = tensor_bytes(chunk_rows, hidden_size, itemsize=hidden_itemsize) - - w1_peak_bytes = silu_bytes + gemm_out_a_bytes - quant_peak_bytes = silu_bytes + quant_bytes + scale_bytes - # A is freed before Q/L. S is freed before B; first-fit can reuse S only - # when B fits there, otherwise B must fit as a contiguous tail allocation. - if silu_bytes >= gemm_out_b_bytes: - temp_bytes = max(w1_peak_bytes, quant_peak_bytes) - else: - temp_bytes = max(w1_peak_bytes, quant_peak_bytes + gemm_out_b_bytes) +) -> int: + """Size one Prefill microbatch workspace for the bounded grouped-MoE path. - max_dense_rows = _align_up( - global_world_size * prefill_tokens_per_rank, - _DEEPEP_GATHER_ROWS_CACHE_ALIGNMENT, - ) - gather_bytes = tensor_bytes(max_dense_rows, hidden_size, itemsize=hidden_itemsize) - per_workspace_bytes = gather_bytes + temp_bytes - prefill_required_total_bytes = per_workspace_bytes * microbatch_count - num_rdma_bytes = _align_up( - max(decode_size_hint, prefill_required_total_bytes), - microbatch_count * _TENSOR_BUFFER_ALIGNMENT_BYTES, - ) - return _LegacyLowLatencyRdmaSizing( - num_rdma_bytes=num_rdma_bytes, - per_workspace_bytes=per_workspace_bytes, - prefill_required_total_bytes=prefill_required_total_bytes, - max_dense_rows=max_dense_rows, - microbatch_count=microbatch_count, + The target chunk covers balanced expert routing in one pass. More skewed + routing remains correct because the consumer already splits expanded rows. + """ + tensor_alignment = 256 + expert_alignment = 128 + metadata_row_granularity = 1024 + + assert num_experts % world_size == 0 + assert intermediate_size % expert_alignment == 0 + + def align_up(value: int, alignment: int) -> int: + return (value + alignment - 1) // alignment * alignment + + hidden_bytes = torch.empty((), dtype=hidden_dtype).element_size() + max_gather_rows = align_up(world_size * num_max_tokens_per_rank, metadata_row_granularity) + num_local_experts = num_experts // world_size + target_chunk_rows = align_up( + num_max_tokens_per_rank * num_experts_per_tok + num_local_experts * (expert_alignment - 1), + expert_alignment, ) + gather_out = align_up(max_gather_rows * hidden_size * hidden_bytes, tensor_alignment) + silu_out = align_up(target_chunk_rows * intermediate_size * hidden_bytes, tensor_alignment) + gemm_out_a = align_up(target_chunk_rows * 2 * intermediate_size * hidden_bytes, tensor_alignment) + quant_out = align_up(target_chunk_rows * intermediate_size, tensor_alignment) + quant_scale = align_up(target_chunk_rows * (intermediate_size // expert_alignment) * 4, tensor_alignment) + gemm_out_b = align_up(target_chunk_rows * hidden_size * hidden_bytes, tensor_alignment) + + # TensorBufferManager uses first-fit allocation. W1 keeps silu_out and gemm_out_a + # together; after gemm_out_a is reused by quantization, the freed silu_out block + # can only hold gemm_out_b when it is large enough. + w1_peak = gather_out + silu_out + gemm_out_a + w2_peak = gather_out + silu_out + quant_out + quant_scale + if gemm_out_b > silu_out: + w2_peak += gemm_out_b + # TensorBufferManager may trim an unaligned prefix from the supplied view. + return max(w1_peak, w2_peak) + tensor_alignment + try: import deep_ep @@ -336,39 +294,27 @@ def new_deepep_group( enable_mega_moe_buffer = False has_legacy_moe_layer = True - # The legacy buffer is also the only prefill workspace. Pure prefill - # does not use its low-latency protocol, but still needs its RDMA - # storage for the bounded grouped-MoE compute path. - enable_low_latency_buffer = has_legacy_moe_layer + enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" + enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" if enable_low_latency_buffer: - if self.ll_decode_num_tokens is None: - decode_size_hint = 0 - else: - decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( - self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts - ) - rdma_sizing = _calculate_legacy_low_latency_rdma_sizing( - decode_size_hint=decode_size_hint, - global_world_size=global_world_size, - prefill_tokens_per_rank=prefill_num_max_dispatch_tokens_per_rank, - hidden_size=hidden_size, - moe_intermediate_size=moe_intermediate_size, - hidden_dtype=get_torch_dtype(args.data_type), - microbatch_count=len(self.groups), - ) - num_rdma_bytes = rdma_sizing.num_rdma_bytes - logger.info( - "Initialize DeepEP legacy RDMA workspace: decode_size_hint=%s, " - "prefill_required_total_bytes=%s, num_rdma_bytes=%s, " - "per_workspace_bytes=%s, max_dense_rows=%s, microbatch_count=%s", - decode_size_hint, - rdma_sizing.prefill_required_total_bytes, - num_rdma_bytes, - rdma_sizing.per_workspace_bytes, - rdma_sizing.max_dense_rows, - rdma_sizing.microbatch_count, + # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 + # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 + num_rdma_bytes = deep_ep.Buffer.get_low_latency_rdma_size_hint( + self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts ) + # normal 节点同时执行 Prefill 和 Decode,复用的 RDMA buffer 必须覆盖全部 Prefill workspace。 + if args.run_mode == "normal": + workspace_size = get_deep_ep_prefill_moe_workspace_size( + num_max_tokens_per_rank=self.ll_num_tokens, + hidden_size=self.ll_hidden, + intermediate_size=moe_intermediate_size, + num_experts_per_tok=num_experts_per_tok, + num_experts=self.ll_num_experts, + world_size=global_world_size, + hidden_dtype=get_torch_dtype(args.data_type), + ) + num_rdma_bytes = max(num_rdma_bytes, workspace_size * len(self.groups)) self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, num_rdma_bytes=num_rdma_bytes, @@ -379,6 +325,22 @@ def new_deepep_group( torch.uint8, use_rdma_buffer=True ) + if enable_prefill_workspace: + workspace_size = get_deep_ep_prefill_moe_workspace_size( + num_max_tokens_per_rank=self.ll_num_tokens, + hidden_size=self.ll_hidden, + intermediate_size=moe_intermediate_size, + num_experts_per_tok=num_experts_per_tok, + num_experts=self.ll_num_experts, + world_size=global_world_size, + hidden_dtype=get_torch_dtype(args.data_type), + ) + self.ep_prefill_moe_workspace = torch.empty( + workspace_size * len(self.groups), + dtype=torch.uint8, + device=torch.device("cuda", torch.cuda.current_device()), + ) + if enable_mega_moe_buffer: # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy # low-latency buffer,因此纯 FP4 模型无需承担后者的大块 RDMA 显存。 @@ -425,12 +387,12 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: """Return a slice of the workspace used by DeepEP Prefill MoE kernels. - All legacy FP8 MoE modes, including pure prefill, use local DeepEP - low-latency RDMA storage. Pure prefill only treats it as workspace; - decode-capable modes reuse the same storage after decode has released - its low-latency layout. With one communication group, the default + Pure Prefill nodes own a dedicated workspace and do not initialize a + low-latency Decode buffer. Decode-capable nodes reuse the idle local + RDMA storage. With one communication group, the default ``microbatch_index=0`` receives the whole workspace. With multiple - groups, the workspace is split into ``len(self.groups)`` equal slices. + groups, the workspace is split into ``len(self.groups)`` equal slices + and each in-flight microbatch uses the slice matching its group index. Args: microbatch_index: Zero-based microbatch and communication-group @@ -454,11 +416,10 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. def clear_deepep_buffer(self): """ Decode-capable modes reuse the low-latency RDMA buffer during Prefill, - so clean it before the next low-latency Decode. Pure Prefill uses the - same RDMA storage only as workspace and has no low-latency layout to - clean. + so clean it before the next low-latency Decode. Pure Prefill owns a + dedicated workspace and has no low-latency buffer to clean. """ - if self.ep_low_latency_buffer is not None and self.ll_decode_num_tokens is not None: + if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts ) From 69c260f4a5f193fcd9170d06d0c77187a2ee18c1 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 24 Aug 2026 05:07:19 +0000 Subject: [PATCH 126/214] set qp_depth minimum 128 --- lightllm/utils/envs_utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index e121e2f193..cb5fc4fcc4 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -52,7 +52,8 @@ def set_env_start_args(args): else: decode_capacity = get_deepep_num_max_dispatch_tokens_per_rank_decode() min_qp_depth = 2 * (decode_capacity + 1) - derived_qp_depth = 1 << (min_qp_depth - 1).bit_length() + # NVSHMEM IBGDA rejects QP depths below NVSHMEMI_IBGDA_MIN_QP_DEPTH. + derived_qp_depth = max(128, 1 << (min_qp_depth - 1).bit_length()) configured_qp_depth = int(os.getenv("NVSHMEM_QP_DEPTH", derived_qp_depth)) if configured_qp_depth < derived_qp_depth: logger.warning( From baadd7ea316861c089956d46021e0b216a51bcf8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:10:32 +0000 Subject: [PATCH 127/214] perf(server): reduce PD streaming and parser overhead - avoid per-message tasks when receiving queued PD packets - scan DSML buffers only when closing tags arrive - validate request lengths directly from token counts - redact inline image data from prompt-build error logs --- lightllm/server/api_http_pd.py | 9 ++-- lightllm/server/build_prompt.py | 22 ++++++++- lightllm/server/function_call_parser.py | 11 ++++- lightllm/server/httpserver/manager.py | 13 ++---- .../httpserver_for_pd_master/manager.py | 46 +++++++++++++++---- 5 files changed, 79 insertions(+), 22 deletions(-) diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 1d8b2112fc..54ff2ab5b5 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -8,10 +8,10 @@ ``g_objs`` 在 handler 内懒导入,避免与 api_http 循环依赖。 """ -import asyncio import pickle import ujson as json +from anyio import fail_after from fastapi import APIRouter, WebSocket, WebSocketDisconnect from lightllm.server.pd_io_struct import ObjType @@ -38,13 +38,16 @@ async def register_and_keep_alive(websocket: WebSocket): try: heartbeat_timeout_seconds = 30 while True: - data = await asyncio.wait_for(websocket.receive_bytes(), timeout=heartbeat_timeout_seconds) + # Avoid creating a new Task per message so queued PD token packs + # can be drained without extra event-loop scheduling. + with fail_after(heartbeat_timeout_seconds): + data = await websocket.receive_bytes() obj = pickle.loads(data) if isinstance(obj, tuple) and obj and obj[0] == ObjType.HEARTBEAT: continue await g_objs.httpserver_manager.put_to_handle_queue(obj) - except asyncio.TimeoutError: + except TimeoutError: logger.warning(f"client {regist_json} heartbeat timed out after {heartbeat_timeout_seconds} seconds") try: await websocket.close(code=1011, reason="PD heartbeat timed out") diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index a28e51dd9e..1ee9994643 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -133,6 +133,23 @@ def _normalize_multimodal_content_types(messages: list) -> None: part["type"] = "audio" +def _redact_message_image_data_in_place(messages: list) -> None: + for message in messages: + content = message.get("content") + if not isinstance(content, list): + continue + + for part in content: + image_url = part.get("image_url") + if not image_url: + continue + + url = image_url["url"] + if url.startswith("data:image"): + media_type, separator, _ = url.partition(",") + image_url["url"] = f"{media_type}{separator}" if separator else "" + + async def build_prompt(request, tools) -> str: # pydantic格式转成dict, 否则,当根据tokenizer_config.json拼template时,Jinja判断无法识别 messages = [m.model_dump(by_alias=True, exclude_none=True) for m in request.messages] @@ -169,9 +186,12 @@ async def build_prompt(request, tools) -> str: try: input_str = tokenizer.apply_chat_template(**kwargs, tokenize=False, add_generation_prompt=True, tools=tools) except Exception as e: + request_dump = request.model_dump(by_alias=True, exclude_none=True) + _redact_message_image_data_in_place(request_dump["messages"]) + _redact_message_image_data_in_place(kwargs["conversation"]) logger.exception( "Failed to build prompt. request=%s tools=%s template_kwargs=%s", - json.dumps(request.model_dump(by_alias=True, exclude_none=True), ensure_ascii=False, default=str), + json.dumps(request_dump, ensure_ascii=False, default=str), json.dumps(tools, ensure_ascii=False, default=str), json.dumps(kwargs, ensure_ascii=False, default=str), ) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index 8ee42ac345..a4bba104af 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1572,6 +1572,10 @@ def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> StreamingParseResult: """Streaming incremental parsing for DSML format tool calls.""" + overlap = max(len(self.invoke_end_token), len(self.param_end_token)) - 1 + new_input_window = self._buffer[-overlap:] + new_text + has_new_invoke_end = self.invoke_end_token in new_input_window + has_new_param_end = self.param_end_token in new_input_window self._buffer += new_text normal_text_parts = [] calls: List[ToolCallItem] = [] @@ -1621,7 +1625,7 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami if self.eot_token.startswith(current_text): return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - complete_invoke_match = self.invoke_regex.match(current_text) + complete_invoke_match = self.invoke_regex.match(current_text) if has_new_invoke_end else None if complete_invoke_match: func_name = complete_invoke_match.group(1) invoke_body = complete_invoke_match.group(2) @@ -1681,6 +1685,9 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami self.streamed_args_for_tool.append("") continue + if self.current_tool_name_sent and not has_new_param_end: + return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + partial_match = self.partial_invoke_regex.match(current_text) if not partial_match: return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) @@ -1712,7 +1719,7 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami "name": func_name, "arguments": {}, } - else: + if has_new_param_end: # Stream arguments as complete parameters are parsed param_matches = self.param_regex.findall(partial_body) if param_matches and len(param_matches) > len(self._accumulated_params): diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 89fd01bcfc..d9ec74af3a 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -396,7 +396,7 @@ async def generate( ) prompt_tokens = len(prompt_ids) - prompt_ids = await self._check_and_repair_length(prompt_ids, sampling_params) + self._check_and_repair_length(prompt_tokens, sampling_params) # 监控 self.metric_client.counter_inc("lightllm_request_count") self.metric_client.histogram_observe("lightllm_request_input_length", prompt_tokens) @@ -626,10 +626,9 @@ def get_real_supported_max_req_total_len(self): # 得到系统真正能支持的最大长度,同时收到启动参数中模型支持长度的限制,也收到token容量的限制。 return min(self.shm_max_total_token_num.get_value() - 36, self.max_req_total_len) - async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params: SamplingParams): - if not prompt_ids: - raise ValueError("prompt_ids is empty") - prompt_tokens = len(prompt_ids) + def _check_and_repair_length(self, prompt_tokens: int, sampling_params: SamplingParams) -> None: + if prompt_tokens <= 0: + raise ValueError("prompt_tokens must be greater than 0") # 这里 -36 是保留一些不可预知的边界余量,防止系统出错 real_supported_max_req_total_len = self.get_real_supported_max_req_total_len() @@ -651,14 +650,12 @@ async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params: ) # last repaired - req_total_len = len(prompt_ids) + sampling_params.max_new_tokens + req_total_len = prompt_tokens + sampling_params.max_new_tokens if req_total_len > self.max_req_total_len: raise ValueError( f"the req total len (input len + output len) is too long > max_req_total_len:{self.max_req_total_len}" ) - return prompt_ids - async def transfer_to_next_module_or_node( self, prompt: str, diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 2f31b35aeb..071c8e24f2 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -54,6 +54,7 @@ def __init__( self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 + self.next_request_queue_metric_time = 0.0 self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) @@ -62,7 +63,7 @@ def __init__( return def get_real_supported_max_req_total_len(self): - # HttpServerManager.generate 会借用 _check_and_repair_length(self, ...),其中会调用本方法。 + # HttpServerManager.generate 会借用长度校验逻辑,其中会调用本方法。 # PD master 无本地 token 池 shm 计数;上限与启动参数及子节点对齐的 max_req_total_len 一致。 return self.max_req_total_len @@ -159,12 +160,9 @@ async def _generate( # 计算输入的 input_token_num, 进行校验,如果输入+输出参数设置太长,则将 # sampling_params 的参数进行修正。 input_token_num = await asyncio.to_thread(self.tokens, prompt, multimodal_params, sampling_params) - fake_prompt_ids = [0 for _ in range(input_token_num)] from lightllm.server.httpserver.manager import HttpServerManager - await HttpServerManager._check_and_repair_length( - self, prompt_ids=fake_prompt_ids, sampling_params=sampling_params - ) + HttpServerManager._check_and_repair_length(self, prompt_tokens=input_token_num, sampling_params=sampling_params) return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) @@ -458,7 +456,14 @@ async def fetch_pd_stream( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) + request_queue_duration = req_status.oldest_age() token_list = req_status.pop_all_tokens() + if token_list and now >= self.next_request_queue_metric_time: + self.metric_client.histogram_observe( + "lightllm_pd_master_request_queue_duration", request_queue_duration + ) + self.next_request_queue_metric_time = now + 1.0 + ready_token_list = [] for sub_req_id, request_output, metadata, finish_status in token_list: output_index = metadata.get("count_output_tokens") # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 @@ -470,12 +475,16 @@ async def fetch_pd_stream( if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - yield sub_req_id, request_output, metadata, finish_status + ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) else: continue else: metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - yield sub_req_id, request_output, metadata, finish_status + ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) + + for index, token_info in enumerate(ready_token_list): + token_info[2]["_pd_stream_batch_end"] = index == len(ready_token_list) - 1 + yield token_info return @@ -624,6 +633,8 @@ async def put_to_handle_queue(self, obj): async def handle_loop(self): self.infos_queues = AsyncQueue() asyncio.create_task(self.timer_log()) + event_loop = asyncio.get_running_loop() + next_ingress_queue_metric_time = 0.0 use_config_server = self.args.config_server_host and self.args.config_server_port @@ -633,7 +644,15 @@ async def handle_loop(self): asyncio.create_task(register_loop(self)) while True: - objs = await self.infos_queues.wait_to_get_all_data() + await self.infos_queues.wait_to_ready() + ingress_queue_duration = self.infos_queues.oldest_age() + objs = await self.infos_queues.get_all_data() + now = event_loop.time() + if objs and now >= next_ingress_queue_metric_time: + self.metric_client.histogram_observe( + "lightllm_pd_master_ingress_queue_duration", ingress_queue_duration + ) + next_ingress_queue_metric_time = now + 1.0 try: for obj in objs: @@ -686,6 +705,7 @@ def __init__(self, req_id, p_node, d_node) -> None: self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() self.out_token_info_list: List[Tuple[int, str, dict, FinishStatus]] = [] + self.oldest_token_time = None self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node @@ -701,19 +721,29 @@ def append_token(self, token_info: Tuple[int, str, dict, FinishStatus]): was_empty = not self.out_token_info_list self.out_token_info_list.append(token_info) if was_empty: + self.oldest_token_time = time.monotonic() self.event.set() + def oldest_age(self): + if self.oldest_token_time is None: + return 0.0 + return time.monotonic() - self.oldest_token_time + def pop_all_tokens(self): self.event.clear() ans = self.out_token_info_list self.out_token_info_list = [] + self.oldest_token_time = None return ans def put_tokens_to_front(self, token_list: List[Tuple[int, str, dict, FinishStatus]]): if not token_list: return + was_empty = not self.out_token_info_list self.out_token_info_list = token_list + self.out_token_info_list + if was_empty: + self.oldest_token_time = time.monotonic() self.event.set() From 095c043d4781c8cc8f8b515d4e4e618958ff5a6f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:20:49 +0000 Subject: [PATCH 128/214] perf(openai): reduce chat streaming serialization overhead - serialize common SSE deltas without Pydantic models - coalesce events at PD stream batch boundaries - cache tool parser and Kimi history metadata per stream - flush buffered events before propagating stream errors --- lightllm/server/api_openai.py | 235 +++++++++++++++++++--------------- 1 file changed, 133 insertions(+), 102 deletions(-) diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index 5ebccb42f0..2f8b8e6712 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -98,6 +98,35 @@ def _serialize_sse_chunk(chunk, choice_nulls=(), response_nulls=()): return json.dumps(d, ensure_ascii=False) +def _serialize_chat_sse_chunk( + request_id, + created, + model, + choice_index, + delta, + choice_nulls=(), + response_nulls=(), + finish_reason=None, +): + """Serialize a chat streaming chunk without constructing Pydantic models.""" + choice = {"index": int(choice_index), "delta": delta} + if finish_reason is not None: + choice["finish_reason"] = finish_reason + for field in choice_nulls: + choice[field] = None + + chunk = { + "id": request_id, + "object": "chat.completion.chunk", + "created": created, + "model": model, + "choices": [choice], + } + for field in response_nulls: + chunk[field] = None + return json.dumps(chunk, ensure_ascii=False) + + def _process_tool_call_id( tool_call_parser, call_item: ToolCallItem, @@ -478,11 +507,52 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req _final_choice_nulls = ("logprobs", "token_ids", "stop_reason") _first_resp_nulls = ("prompt_token_ids",) + def make_chat_sse_chunk(choice_index, delta, choice_nulls=(), response_nulls=(), finish_reason=None): + payload = _serialize_chat_sse_chunk( + chat_completion_id, + created_time, + request.model, + choice_index, + delta, + choice_nulls, + response_nulls, + finish_reason, + ) + return f"data: {payload}\n\n" + + pending_sse_chunks = [] + pending_sse_chars = 0 + + def append_sse_chunk(chunk): + nonlocal pending_sse_chars + pending_sse_chunks.append(chunk) + pending_sse_chars += len(chunk) + + def pending_sse_limit_reached(): + return len(pending_sse_chunks) >= 32 or pending_sse_chars >= 64 * 1024 + + def pop_sse_chunk(): + nonlocal pending_sse_chars + chunk_count = 0 + chunk_chars = 0 + for pending_chunk in pending_sse_chunks: + if chunk_count and (chunk_count >= 32 or chunk_chars + len(pending_chunk) > 64 * 1024): + break + chunk_count += 1 + chunk_chars += len(pending_chunk) + + chunk = "".join(pending_sse_chunks[:chunk_count]) + del pending_sse_chunks[:chunk_count] + pending_sse_chars -= chunk_chars + return chunk + # Streaming case - async def stream_results() -> AsyncGenerator[bytes, None]: + async def stream_results_inner() -> AsyncGenerator[bytes, None]: has_emitted_tool_calls: Dict[int, bool] = collections.defaultdict(bool) has_emitted_first_chunk: Dict[int, bool] = collections.defaultdict(bool) stream_tool_call_ids: Dict[Tuple[int, int], str] = {} + tool_parser = getattr(g_objs.args, "tool_call_parser", None) or "llama3" + history_tool_calls_cnt = _get_history_tool_calls_cnt(request) if tool_parser == "kimi_k2" else 0 from .req_id_generator import convert_sub_id_to_group_id prompt_tokens = 0 @@ -494,6 +564,7 @@ async def stream_results() -> AsyncGenerator[bytes, None]: completion_tokens += 1 group_request_id = convert_sub_id_to_group_id(sub_req_id) choice_index = sub_req_id - group_request_id + pd_stream_batch_end = metadata.pop("_pd_stream_batch_end", True) delta = request_output current_finish_reason = finish_status.get_finish_reason() @@ -502,18 +573,14 @@ async def stream_results() -> AsyncGenerator[bytes, None]: # OpenAI SSE spec: role appears only in the first delta with content="". if not has_emitted_first_chunk[choice_index]: has_emitted_first_chunk[choice_index] = True - first_choice = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(role="assistant", content=""), - finish_reason=None, - ) - first_chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[first_choice], + append_sse_chunk( + make_chat_sse_chunk( + choice_index, + {"role": "assistant", "content": ""}, + _first_choice_nulls, + _first_resp_nulls, + ) ) - yield f"data: {_serialize_sse_chunk(first_chunk, _first_choice_nulls, _first_resp_nulls)}\n\n" # Handle reasoning content if get_env_start_args().reasoning_parser: @@ -522,18 +589,9 @@ async def stream_results() -> AsyncGenerator[bytes, None]: ) if reasoning_text: if request.separate_reasoning: - choice_data = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(reasoning=reasoning_text), - finish_reason=None, + append_sse_chunk( + make_chat_sse_chunk(choice_index, {"reasoning": reasoning_text}, _choice_nulls) ) - chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - choices=[choice_data], - model=request.model, - ) - yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" else: delta = reasoning_text + (delta or "") @@ -545,22 +603,11 @@ async def stream_results() -> AsyncGenerator[bytes, None]: # 1) if there's normal_text, output it as normal content if normal_text and (normal_text.strip() or not has_emitted_tool_calls[sub_req_id]): - choice_data = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(content=normal_text), - finish_reason=None, - ) - chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - choices=[choice_data], - model=request.model, - ) - yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" + append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": normal_text}, _choice_nulls)) # 2) if we found calls, we output them as separate chunk(s) - history_tool_calls_cnt = _get_history_tool_calls_cnt(request) - fc_parser = parser_dict[choice_index] + if calls: + fc_parser = parser_dict[choice_index] for call_item in calls: has_emitted_tool_calls[sub_req_id] = True # transform call_item -> FunctionResponse + ToolCall @@ -582,7 +629,6 @@ async def stream_results() -> AsyncGenerator[bytes, None]: remaining_call = expected_call.replace(actual_call, "", 1) call_item.parameters = remaining_call - tool_parser = getattr(g_objs.args, "tool_call_parser", None) or "llama3" stream_index = getattr(call_item, "tool_index", None) id_key = (choice_index, stream_index) if call_item.name: @@ -619,7 +665,7 @@ async def stream_results() -> AsyncGenerator[bytes, None]: choices=[head_choice], model=request.model, ) - yield f"data: {_serialize_sse_chunk(head_chunk, _choice_nulls)}\n\n" + append_sse_chunk(f"data: {_serialize_sse_chunk(head_chunk, _choice_nulls)}\n\n") for arg_delta in _split_tool_argument_delta(call_item.parameters): arg_tool_call = ToolCall( @@ -637,7 +683,7 @@ async def stream_results() -> AsyncGenerator[bytes, None]: choices=[arg_choice], model=request.model, ) - yield f"data: {_serialize_sse_chunk(arg_chunk, _choice_nulls)}\n\n" + append_sse_chunk(f"data: {_serialize_sse_chunk(arg_chunk, _choice_nulls)}\n\n") else: tool_call = ToolCall( id=tool_call_id if is_tool_head else None, @@ -663,38 +709,31 @@ async def stream_results() -> AsyncGenerator[bytes, None]: choices=[choice_data], model=request.model, ) - yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" + append_sse_chunk(f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n") else: if delta: # If this is the final token, merge content with finish_reason if current_finish_reason is not None: if has_emitted_tool_calls[sub_req_id] and current_finish_reason == "stop": current_finish_reason = "tool_calls" - delta_message = DeltaMessage(content=delta) - stream_choice = ChatCompletionStreamResponseChoice( - index=choice_index, delta=delta_message, finish_reason=current_finish_reason - ) - stream_resp = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[stream_choice], + append_sse_chunk( + make_chat_sse_chunk( + choice_index, + {"content": delta}, + _final_choice_nulls, + finish_reason=current_finish_reason, + ) ) - yield f"data: {_serialize_sse_chunk(stream_resp, _final_choice_nulls)}\n\n" + if pd_stream_batch_end: + while pending_sse_chunks: + yield pop_sse_chunk() + else: + while pending_sse_limit_reached(): + yield pop_sse_chunk() # Skip the separate final-chunk logic below continue else: - delta_message = DeltaMessage(content=delta) - stream_choice = ChatCompletionStreamResponseChoice( - index=choice_index, delta=delta_message, finish_reason=None - ) - stream_resp = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[stream_choice], - ) - yield f"data: {_serialize_sse_chunk(stream_resp, _choice_nulls)}\n\n" + append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": delta}, _choice_nulls)) # Emit a per-choice final chunk with finish_reason (for tool_calls path # or when no delta was emitted alongside finish_reason). @@ -707,52 +746,33 @@ async def stream_results() -> AsyncGenerator[bytes, None]: flush_reasoning, flush_text = parser.flush() if flush_reasoning: if request.separate_reasoning: - flush_choice = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(reasoning=flush_reasoning), - finish_reason=None, - ) + flush_delta = {"reasoning": flush_reasoning} else: # vLLM compat: emit buffered thinking as content - flush_choice = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(content=flush_reasoning), - finish_reason=None, - ) - flush_chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[flush_choice], - ) - yield f"data: {_serialize_sse_chunk(flush_chunk, _choice_nulls)}\n\n" + flush_delta = {"content": flush_reasoning} + append_sse_chunk(make_chat_sse_chunk(choice_index, flush_delta, _choice_nulls)) if flush_text: - flush_choice = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(content=flush_text), - finish_reason=None, - ) - flush_chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[flush_choice], - ) - yield f"data: {_serialize_sse_chunk(flush_chunk, _choice_nulls)}\n\n" + append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": flush_text}, _choice_nulls)) if has_emitted_tool_calls[sub_req_id] and current_finish_reason == "stop": current_finish_reason = "tool_calls" - final_choice = ChatCompletionStreamResponseChoice( - index=choice_index, - delta=DeltaMessage(), - finish_reason=current_finish_reason, - ) - final_chunk = ChatCompletionStreamResponse( - id=chat_completion_id, - created=created_time, - model=request.model, - choices=[final_choice], + append_sse_chunk( + make_chat_sse_chunk( + choice_index, + {}, + _final_choice_nulls, + finish_reason=current_finish_reason, + ) ) - yield f"data: {_serialize_sse_chunk(final_chunk, _final_choice_nulls)}\n\n" + + if pd_stream_batch_end: + while pending_sse_chunks: + yield pop_sse_chunk() + else: + while pending_sse_limit_reached(): + yield pop_sse_chunk() + + while pending_sse_chunks: + yield pop_sse_chunk() usage = UsageInfo( prompt_tokens=prompt_tokens, @@ -771,6 +791,17 @@ async def stream_results() -> AsyncGenerator[bytes, None]: yield "data: [DONE]\n\n".encode("utf-8") + async def stream_results() -> AsyncGenerator[bytes, None]: + try: + async for chunk in stream_results_inner(): + yield chunk + except ClientDisconnected: + raise + except Exception: + while pending_sse_chunks: + yield pop_sse_chunk() + raise + background_tasks = BackgroundTasks() return CustomStreamingResponse( _safe_stream_wrapper(stream_results()), media_type="text/event-stream", background=background_tasks From 5e491fb5db02d01a5476dad4a35be9f722d7c1fb Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:21:29 +0000 Subject: [PATCH 129/214] feat(metrics): add PD pipeline latency metrics --- lightllm/server/api_stream_obj.py | 31 +++++++++++++++++++++++ lightllm/server/httpserver/async_queue.py | 9 +++++++ lightllm/server/httpserver/pd_loop.py | 14 +++++++++- lightllm/server/metrics/metrics.py | 10 ++++++++ 4 files changed, 63 insertions(+), 1 deletion(-) diff --git a/lightllm/server/api_stream_obj.py b/lightllm/server/api_stream_obj.py index f4765a53fd..4e59692621 100644 --- a/lightllm/server/api_stream_obj.py +++ b/lightllm/server/api_stream_obj.py @@ -21,12 +21,41 @@ and keep their existing response-header timeout semantics. """ +import time + from fastapi.responses import StreamingResponse from starlette.types import Send from lightllm.utils.envs_utils import get_env_start_args +_pd_send_next_metric_time = 0.0 +_pd_send_max_duration = 0.0 +_pd_send_max_bytes = 0 + + +def _record_pd_send_metrics(duration, body_size): + global _pd_send_next_metric_time + global _pd_send_max_duration + global _pd_send_max_bytes + + _pd_send_max_duration = max(_pd_send_max_duration, duration) + _pd_send_max_bytes = max(_pd_send_max_bytes, body_size) + now = time.monotonic() + if now < _pd_send_next_metric_time: + return + + from lightllm.server.api_http import g_objs + + g_objs.httpserver_manager.metric_client.histogram_observe( + "lightllm_pd_master_http_send_duration", _pd_send_max_duration + ) + g_objs.httpserver_manager.metric_client.gauge_set("lightllm_pd_master_http_send_bytes", _pd_send_max_bytes) + _pd_send_max_duration = 0.0 + _pd_send_max_bytes = 0 + _pd_send_next_metric_time = now + 1.0 + + class CustomStreamingResponse(StreamingResponse): """Send the PD-master HTTP status only after the first body chunk is ready. @@ -51,7 +80,9 @@ async def stream_response(self, send: Send) -> None: async def send_chunk(chunk): if not isinstance(chunk, (bytes, memoryview)): chunk = chunk.encode(self.charset) + send_start = time.monotonic() await send({"type": "http.response.body", "body": chunk, "more_body": True}) + _record_pd_send_metrics(time.monotonic() - send_start, len(chunk)) async def send_response_start(): # Read status and headers at send time. The first body iteration diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index 47cfed4c88..10b1e18ed1 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -1,10 +1,12 @@ import asyncio +import time class AsyncQueue: def __init__(self): self.datas = [] self.event = asyncio.Event() + self.oldest_put_time = None async def wait_to_ready(self): try: @@ -16,15 +18,22 @@ async def get_all_data(self): self.event.clear() ans = self.datas self.datas = [] + self.oldest_put_time = None return ans async def put(self, obj): was_empty = not self.datas self.datas.append(obj) if was_empty: + self.oldest_put_time = time.monotonic() self.event.set() return + def oldest_age(self): + if self.oldest_put_time is None: + return 0.0 + return time.monotonic() - self.oldest_put_time + async def wait_to_get_all_data(self): await self.wait_to_ready() handle_list = await self.get_all_data() diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 86462e28b8..bcf1ab1165 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -267,11 +267,23 @@ async def _pd_process_generate( # 转发token的task async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): max_message_size = get_lightllm_websocket_max_message_size() + event_loop = asyncio.get_running_loop() + next_queue_metric_time = 0.0 while True: - handle_list = await forwarding_queue.wait_to_get_all_data() + await forwarding_queue.wait_to_ready() + queue_duration = forwarding_queue.oldest_age() + handle_list = await forwarding_queue.get_all_data() if handle_list: + now = event_loop.time() + if now >= next_queue_metric_time: + from lightllm.server.api_http import g_objs + + g_objs.httpserver_manager.metric_client.histogram_observe( + "lightllm_pd_forward_queue_duration", queue_duration + ) + next_queue_metric_time = now + 1.0 load_info: dict = _get_load_info() pending_handle_lists = [] group_start = 0 diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index c19d756c17..b0482f93f6 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -27,6 +27,11 @@ "lightllm_cache_ratio": "cache length / input_length", "lightllm_batch_current_max_tokens": "dynamic max token used for current batch", "lightllm_request_mtp_avg_token_per_step": "Average number of tokens per step", + "lightllm_pd_forward_queue_duration": "Time tokens wait on a P/D node before websocket forwarding (s)", + "lightllm_pd_master_ingress_queue_duration": "Time token packets wait in the PD master ingress queue (s)", + "lightllm_pd_master_request_queue_duration": "Time tokens wait for their PD master request consumer (s)", + "lightllm_pd_master_http_send_duration": "Maximum ASGI body send duration in the latest sample window (s)", + "lightllm_pd_master_http_send_bytes": "Largest ASGI body sent in the latest sample window (bytes)", "lightllm_prompt_tokens_total": "Total number of prefill tokens processed", "lightllm_generation_tokens_total": "Total number of generation tokens processed", "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", @@ -97,10 +102,15 @@ def init_metrics(self, args): self.create_histogram("lightllm_request_mean_time_per_token_duration", self.duration_buckets) self.create_histogram("lightllm_request_first_token_duration", self.duration_buckets) self.create_histogram("lightllm_request_queue_duration_bucket", self.duration_buckets) + self.create_histogram("lightllm_pd_forward_queue_duration", self.duration_buckets) + self.create_histogram("lightllm_pd_master_ingress_queue_duration", self.duration_buckets) + self.create_histogram("lightllm_pd_master_request_queue_duration", self.duration_buckets) + self.create_histogram("lightllm_pd_master_http_send_duration", self.duration_buckets) self.create_histogram("lightllm_batch_inference_duration_bucket", self.duration_buckets, labelnames=["method"]) self.gateway_url = args.metric_gateway self.create_gauge("lightllm_queue_size") + self.create_gauge("lightllm_pd_master_http_send_bytes") self.create_gauge("lightllm_batch_current_size") self.create_gauge("lightllm_batch_pause_size") self.create_gauge("lightllm_batch_current_max_tokens") From f2e8f1429dccd199750bc5d9334dc4a30622dd15 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:50:24 +0000 Subject: [PATCH 130/214] allocate model request state per DP --- lightllm/server/router/manager.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index 01634d962e..32b499b182 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -148,7 +148,11 @@ async def wait_to_model_ready(self): "weight_dir": self.model_weightdir, "load_way": self.load_way, "max_total_token_num": self.max_total_token_num, - "max_req_num": self.args.running_max_req_size, + "max_req_num": ( + self.args.running_max_req_size + if self.args.run_mode in ["prefill", "normal"] and self.args.enable_dp_prompt_cache_fetch + else self.args.per_dp_running_max_req_size + ), "max_seq_length": self.args.max_req_total_len + 8, # 留一点余量 "nccl_host": self.args.nccl_host, "nccl_port": get_shm_port_args().nccl_port, From 7fad48a98421ccb326d529029bc5fa496c3512ce Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:54:41 +0000 Subject: [PATCH 131/214] fix(pd): prevent PD master hang on prefill timeout race --- lightllm/server/httpserver/manager.py | 6 +----- lightllm/server/httpserver_for_pd_master/manager.py | 11 ++++++++++- 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index d9ec74af3a..2535ce5435 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -417,11 +417,7 @@ async def generate( await pd_upload_websocket.send( pickle.dumps((ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS, group_request_id, array("i", prompt_ids))) ) - try: - await asyncio.wait_for(pd_event.wait(), timeout=180) - except asyncio.TimeoutError: - logger.error(f"pd prefill node wait pd_event 180s time out, group_req_id {group_request_id}") - raise Exception(f"group_req_id {group_request_id} wait pd_event time out") + await pd_event.wait() decode_node_info: PDDecodeNodeInfo = pd_event.decode_node_info sampling_params.pd_kv_trans_params.set(pickle.dumps(decode_node_info)) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 071c8e24f2..c047ba75b0 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -468,14 +468,20 @@ async def fetch_pd_stream( output_index = metadata.get("count_output_tokens") # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 if output_index == 1: + node_run_mode = metadata.pop("node_mode", None) if first_token_gen is False: first_token_gen = True - node_run_mode = metadata.pop("node_mode", None) if node_run_mode == "prefill": if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) metadata["prompt_cache_len"] = prompt_cache_len_from_prefill ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) + elif finish_status.status in ( + FinishStatus.FINISHED_ABORTED, + FinishStatus.FINISHED_ERROR, + ): + metadata["prompt_cache_len"] = prompt_cache_len_from_prefill + ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) else: continue else: @@ -519,6 +525,9 @@ async def _wait_for_prefill_token_if_needed( prompt_cache_len = metadata.get("prompt_cache_len", 0) req_status.put_tokens_to_front(new_tokens) return prompt_cache_len + if token[3].is_finished(): + req_status.put_tokens_to_front(new_tokens) + return ready_kv_len async def _wait_to_token_package( self, From 1e09ad1005a1dc58fdbb180955323672d59de601 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 06:58:09 +0000 Subject: [PATCH 132/214] os.getenv(DSV4_SWA_FULL_TOKENS_RATIO) --- lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 6e0cb51e79..305457d166 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,3 +1,5 @@ +import os + import torch from dataclasses import dataclass from typing import List, Optional, Sequence, Union @@ -38,7 +40,7 @@ DSV4_C128_STATE_RING = 128 # 128 rows/request before MTP padding # swa 池占 full token 空间的比例(sglang DSV4 默认 swa_full_tokens_ratio=0.1 同值)。 # 瞬时借页/驱逐走 swa 压力阀;池子大小仅按 ratio 切分,不再叠加结构性余量。 -DSV4_SWA_FULL_TOKENS_RATIO = 0.1 # 0.1 +DSV4_SWA_FULL_TOKENS_RATIO = float(os.getenv("DSV4_SWA_FULL_TOKENS_RATIO", "0.1")) def _ceil_div(a: int, b: int) -> int: From 7dd1ba0149d42f160be25ff63c8baa882a750086 Mon Sep 17 00:00:00 2001 From: sufubao <47234901+sufubao@users.noreply.github.com> Date: Mon, 24 Aug 2026 15:45:07 +0800 Subject: [PATCH 133/214] fix(pd): recover NIXL transfer task failures (#1474) --- .../decode_node_impl/decode_trans_process.py | 3 +- .../prefill_trans_process.py | 93 ++++++++++++------- lightllm/utils/process_check.py | 15 +++ 3 files changed, 78 insertions(+), 33 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 9262c76eb3..edcd28d5d8 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -21,7 +21,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -47,6 +47,7 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 67fbd74cd6..40da12f3a4 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -15,7 +15,7 @@ from ..kv_transporter import create_kv_transporter from lightllm.utils.error_utils import log_exception from lightllm.utils.envs_utils import get_unique_server_name -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread logger = init_logger(__name__) @@ -41,6 +41,7 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): + install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -336,48 +337,76 @@ def update_task_status_loop( time.sleep(0.001) continue - with self.waiting_dict_lock: - tasks = list(self.waiting_dict.values()) - for trans_task in tasks: - if trans_task.xfer_handle is None: - continue + self._update_task_status_once() + time.sleep(0.001) - # 传输任务状态检查 + def _update_task_status_once(self): + with self.waiting_dict_lock: + tasks = list(self.waiting_dict.values()) + for trans_task in tasks: + if trans_task.xfer_handle is None: + continue + + # A native NIXL error must fail this task without terminating the + # status thread and stranding every transfer queued behind it. + try: ret = self.transporter.check_task_status(trans_task=trans_task) - if ret == "DONE": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - if self.transporter.capture_telemetry: - telem = self.transporter.nixl_agent.get_xfer_telemetry(trans_task.xfer_handle) + except BaseException as e: + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + logger.exception(f"check transfer status failed: {failed_task.to_str()}") + failed_task.error_info = f"check transfer status failed: {str(e)}" + self.failed_queue.put(failed_task) + continue + + if ret == "DONE": + completed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if completed_task is None: + continue + if self.transporter.capture_telemetry: + try: + telem = self.transporter.nixl_agent.get_xfer_telemetry(completed_task.xfer_handle) total_us = telem.xferDuration post_us = telem.postDuration backend_us = telem.xferDuration - telem.postDuration - nixl_backend = self.transporter.nixl_agent.query_xfer_backend(trans_task.xfer_handle) + nixl_backend = self.transporter.nixl_agent.query_xfer_backend(completed_task.xfer_handle) logger.info( - f"write trans task request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index} " + f"write trans task request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} " + f"dst_page={completed_task.dst_page_index} " f"xfer time: {total_us:.3f} us, " f"post time: {post_us:.3f} us, backend time: {backend_us:.3f} us, " f"nixl_backend: {nixl_backend}, total_bytes: {telem.totalBytes}" ) - self.transporter.send_write_done_task_to_decode_node(trans_task) - logger.info( - f"send WRITE done nixl notify " - f"request_id={trans_task.request_id} " - f"kv=[{trans_task.start_kv_index},{trans_task.end_kv_index}) " - f"src_page={trans_task.src_page_index} dst_page={trans_task.dst_page_index}" - ) - self.success_queue.put(trans_task) - elif ret == "ERR": - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "xfer error" - self.failed_queue.put(trans_task) - elif trans_task.time_out(): - trans_task = self.waiting_dict.pop(trans_task.get_key(), None) - trans_task.error_info = "time out in update_task_status_loop" - self.failed_queue.put(trans_task) + except BaseException: + logger.exception(f"get transfer telemetry failed: {completed_task.to_str()}") + + try: + self.transporter.send_write_done_task_to_decode_node(completed_task) + except BaseException as e: + logger.exception(f"send WRITE done nixl notify failed: {completed_task.to_str()}") + completed_task.error_info = f"send WRITE done nixl notify failed: {str(e)}" + self.failed_queue.put(completed_task) + continue - time.sleep(0.001) + logger.info( + f"send WRITE done nixl notify " + f"request_id={completed_task.request_id} " + f"kv=[{completed_task.start_kv_index},{completed_task.end_kv_index}) " + f"src_page={completed_task.src_page_index} dst_page={completed_task.dst_page_index}" + ) + self.success_queue.put(completed_task) + elif ret == "ERR": + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "xfer error" + self.failed_queue.put(failed_task) + elif trans_task.time_out(): + failed_task = self.waiting_dict.pop(trans_task.get_key(), None) + if failed_task is not None: + failed_task.error_info = "time out in update_task_status_loop" + self.failed_queue.put(failed_task) @log_exception def success_loop(self): diff --git a/lightllm/utils/process_check.py b/lightllm/utils/process_check.py index 00cc258bf4..db508c8a6c 100644 --- a/lightllm/utils/process_check.py +++ b/lightllm/utils/process_check.py @@ -8,6 +8,21 @@ logger = init_logger(__name__) +def install_fatal_thread_excepthook(): + """Terminate the process when an unexpected daemon-thread exception escapes.""" + + def exit_process(args): + try: + logger.error( + f"fatal background thread failure in {getattr(args.thread, 'name', 'unknown')}", + exc_info=(args.exc_type, args.exc_value, args.exc_traceback), + ) + finally: + os._exit(1) + + threading.excepthook = exit_process + + def is_process_active(pid): try: process = psutil.Process(pid) From 25c0e573c75abfae81d360dc7d9f1a7e1ec6be17 Mon Sep 17 00:00:00 2001 From: sufubao <47234901+sufubao@users.noreply.github.com> Date: Mon, 24 Aug 2026 16:37:10 +0800 Subject: [PATCH 134/214] fix: detach request shared-memory handles on release (#1482) --- lightllm/server/core/objs/req.py | 8 ++++++++ lightllm/server/core/objs/shm_req_manager.py | 2 ++ unit_tests/server/core/objs/test_req.py | 12 ++++++++++++ unit_tests/server/core/objs/test_shm_req_manager.py | 10 ++++++++++ 4 files changed, 32 insertions(+) diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b268120c90..2ebc0ab238 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -312,6 +312,14 @@ def link_logprobs_shm_array(self): self.shm_logprobs.link_shm() return + def detach_shm_arrays(self): + """Detach process-local request-scoped SHM handles before slot reuse.""" + for attr_name in ("shm_prompt_ids", "shm_logprobs"): + shm_array = getattr(self, attr_name, None) + if shm_array is not None: + shm_array.detach_shm() + delattr(self, attr_name) + async def merge_final_token_metadata( self, metadata: Dict[str, Any], diff --git a/lightllm/server/core/objs/shm_req_manager.py b/lightllm/server/core/objs/shm_req_manager.py index fd9106d59c..01d81fde05 100644 --- a/lightllm/server/core/objs/shm_req_manager.py +++ b/lightllm/server/core/objs/shm_req_manager.py @@ -136,6 +136,8 @@ def put_back_req_obj(self, req: Req): req_index_in_mem = req.index_in_shm_mem assert req_index_in_mem < self.max_req_num assert self.proc_private_get_state[req_index_in_mem] == 1 + req.detach_shm_arrays() + with self.get_req_lock_by_index(req_index_in_mem): req.ref_count = req.ref_count - 1 self.proc_private_get_state[req_index_in_mem] = 0 diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index 584f421928..966221e6c4 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -37,6 +37,18 @@ def test_create_prompt_ids_shm_array(req): assert hasattr(req, "shm_prompt_ids") +def test_detach_shm_arrays(req): + prompt_ids = req.shm_prompt_ids + logprobs = req.shm_logprobs + + req.detach_shm_arrays() + + assert prompt_ids.shm is None + assert logprobs.shm is None + assert not hasattr(req, "shm_prompt_ids") + assert not hasattr(req, "shm_logprobs") + + def test_get_used_tokens(req): req.shm_cur_kv_len = 5 assert req.get_used_tokens() == 5 diff --git a/unit_tests/server/core/objs/test_shm_req_manager.py b/unit_tests/server/core/objs/test_shm_req_manager.py index e26f128d5b..814df20c7f 100644 --- a/unit_tests/server/core/objs/test_shm_req_manager.py +++ b/unit_tests/server/core/objs/test_shm_req_manager.py @@ -1,6 +1,8 @@ import os import pytest import time +from unittest.mock import MagicMock + from easydict import EasyDict from lightllm.utils.envs_utils import set_env_start_args, get_env_start_args from lightllm.server.core.objs.shm_req_manager import ShmReqManager @@ -67,8 +69,16 @@ def test_get_req_obj_by_index(shm_req_manager): def test_put_back_req_obj(shm_req_manager): index = shm_req_manager.alloc_req_index() req_obj = shm_req_manager.get_req_obj_by_index(index) + prompt_ids = req_obj.shm_prompt_ids = MagicMock() + logprobs = req_obj.shm_logprobs = MagicMock() + shm_req_manager.put_back_req_obj(req_obj) + assert req_obj.ref_count == 0 + prompt_ids.detach_shm.assert_called_once_with() + logprobs.detach_shm.assert_called_once_with() + assert not hasattr(req_obj, "shm_prompt_ids") + assert not hasattr(req_obj, "shm_logprobs") shm_req_manager.release_req_index(index) From 2b2e8e7627b4beb3c5639ca1d1d68bd363092e5d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 09:36:42 +0000 Subject: [PATCH 135/214] fix(pd): propagate node generation failures (#1494) Report P/D generation errors to the PD master and wake blocked requests so they abort promptly instead of waiting for stage timeouts. Preserve support_ds4's compact-token protocol by assigning the error message ObjType value 8, and delay prefill health accounting until decode assignment completes. --- lightllm/server/httpserver/manager.py | 29 +++++++++++++---- lightllm/server/httpserver/pd_loop.py | 12 ++++++- .../httpserver_for_pd_master/manager.py | 32 +++++++++++++++++++ lightllm/server/pd_io_struct.py | 1 + 4 files changed, 66 insertions(+), 8 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 2535ce5435..e3abd83ca2 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -354,11 +354,10 @@ async def generate( image_count=image_count, ) - async with self._run_reqs_count_lock: - prev = self.run_reqs_count_mark.get_value() - self.run_reqs_count_mark.set_value(prev + 1) - if prev == 0: - self.latest_success_infer_time_mark.set_value(int(time.time())) + running_request_registered = False + if not self.pd_mode.is_P(): + await self._register_running_request() + running_request_registered = True try: # RL:进入 generation admission。若当前处于 pause_generation / abort, @@ -427,6 +426,11 @@ async def generate( # 直接 raise PDPrefillNodeStopGenToken raise PDPrefillNodeStopGenToken(group_request_id=group_request_id) + if self.pd_mode.is_P(): + # 等待 decode 分配期间尚未进入本地推理,不应触发 prefill 推理健康超时。 + await self._register_running_request() + running_request_registered = True + # 申请资源并存储 alloced_req_indexes = [] while len(alloced_req_indexes) < sampling_params.n: @@ -526,8 +530,8 @@ async def generate( # 防止 pending 请求泄漏导致 pause 无法正确结束。 if self.rl_controller is not None: await self.rl_controller.unregister_generation_admission(group_request_id) - async with self._run_reqs_count_lock: - self.run_reqs_count_mark.set_value(self.run_reqs_count_mark.get_value() - 1) + if running_request_registered: + await self._unregister_running_request() return def _count_multimodal_tokens(self, multimodal_params: MultimodalParams) -> Tuple[int, int]: @@ -1000,6 +1004,17 @@ async def handle_loop(self): self.recycle_event.set() return + async def _register_running_request(self): + async with self._run_reqs_count_lock: + prev = self.run_reqs_count_mark.get_value() + self.run_reqs_count_mark.set_value(prev + 1) + if prev == 0: + self.latest_success_infer_time_mark.set_value(int(time.time())) + + async def _unregister_running_request(self): + async with self._run_reqs_count_lock: + self.run_reqs_count_mark.set_value(self.run_reqs_count_mark.get_value() - 1) + class ReqStatus: def __init__(self, group_request_id, multimodal_params, req_objs: List[Req], start_time) -> None: diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index bcf1ab1165..7030880acf 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -258,8 +258,18 @@ async def _pd_process_generate( await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") + except asyncio.CancelledError: + # PD master 主动 abort 或连接断开时不需要反向上报生成错误。 + pass except BaseException as e: - logger.error(str(e)) + group_request_id = sampling_params.group_request_id + logger.exception(f"pd node generate request {group_request_id} failed: {str(e)}") + try: + await pd_upload_websocket.send( + pickle.dumps((ObjType.PD_UPLOAD_GENERATE_ERROR, group_request_id, f"{type(e).__name__}: {str(e)}")) + ) + except Exception: + logger.exception(f"report pd node generate error failed, group_request_id: {group_request_id}") finally: manager.cancel_pd_request_registration(sampling_params.group_request_id) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index c047ba75b0..7c87966eab 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -406,6 +406,7 @@ async def fetch_pd_stream( except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait prefill prompt ids time out") raise + req_status.raise_if_error() prompt_ids = prefill_prompt_ids_event.prompt_ids logger.info(f"group_request_id: {group_request_id} get prefill prompt ids len {len(prompt_ids)}") @@ -426,6 +427,7 @@ async def fetch_pd_stream( except ServerBusyError: logger.warning(f"group_request_id: {group_request_id} wait decode stage time out err, server is busy now.") raise + req_status.raise_if_error() # 将 decode 节点上报的当前请求使用的decode节点的信息下发给 p 节点,这样 p 节点才知道将 kv 传输给那个 d 节点。 upkv_status: PDUpKVStatus = up_status_event.upkv_status @@ -448,6 +450,7 @@ async def fetch_pd_stream( next_disconnect_check = 0.0 while True: await req_status.wait_to_ready() + req_status.raise_if_error() now = time.monotonic() if now >= next_disconnect_check: next_disconnect_check = now + 1.0 @@ -508,6 +511,7 @@ async def _wait_for_prefill_token_if_needed( new_tokens = [] while True: await req_status.wait_to_ready() + req_status.raise_if_error() if await request.is_disconnected(): raise ClientDisconnected( group_request_id=group_request_id, @@ -692,6 +696,18 @@ async def handle_loop(self): logger.error( f"PD_UPLOAD_PREFILL_PROMPT_IDS fail find req status for group_req_id: {group_req_id}" ) + elif obj[0] == ObjType.PD_UPLOAD_GENERATE_ERROR: + _, group_req_id, error_info = obj + logger.error( + f"received PD node generate error, group_req_id: {group_req_id}, error: {error_info}" + ) + req_status = self.req_id_to_out_inf.get(group_req_id) + if req_status is None: + logger.error( + f"PD_UPLOAD_GENERATE_ERROR fail find req status for group_req_id: {group_req_id}" + ) + else: + req_status.set_error(error_info) else: logger.error(f"recevie error obj {obj}") except BaseException as e: @@ -717,6 +733,7 @@ def __init__(self, req_id, p_node, d_node) -> None: self.oldest_token_time = None self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node + self.error_info: Optional[str] = None async def wait_to_ready(self): try: @@ -724,6 +741,21 @@ async def wait_to_ready(self): except asyncio.TimeoutError: pass + def set_error(self, error_info: str): + # handle_loop and request consumers run on the same event loop. + self.error_info = error_info + self.event.set() + self.up_status_event.set() + self.prefill_prompt_ids_event.set() + + def raise_if_error(self): + if self.error_info is not None: + logger.error( + f"group_request_id: {self.req_id} detected PD node generate error, " + f"raise exception to end the request flow early: {self.error_info}" + ) + raise RuntimeError(f"PD node generate failed: {self.error_info}") + def append_token(self, token_info: Tuple[int, str, dict, FinishStatus]): # TOKEN_PACKS handling and fetch_pd_stream run on the same event loop. Keeping # the mutation free of awaits makes the empty -> ready transition atomic. diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index 36189309a3..da9d53522e 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -42,6 +42,7 @@ class ObjType(enum.Enum): PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 TOKEN_PACKS_COMPACT = 7 # 不含 logprobs 等可选字段的紧凑 token 包。 + PD_UPLOAD_GENERATE_ERROR = 8 # P/D 节点向 pd master 上报本地请求生成异常。 PD_COMPACT_TOKEN_INFO_LEN = 9 From d5b04bf02864d1975880400379392efab13345b7 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 25 Aug 2026 14:46:20 +0000 Subject: [PATCH 136/214] fix(dp): serialize DP collectives after overlap forward The next infer thread may launch DP all-gathers before the previous overlap-stream forward finishes, causing NCCL to overlap with DeepEP/MTP kernels. Make the current stream wait before starting DP collectives while preserving overlap for request preparation. --- .../router/model_infer/mode_backend/dp_backend/impl.py | 6 ++++++ 1 file changed, 6 insertions(+) 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 834b97fa8d..b5256fa967 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 @@ -112,6 +112,12 @@ def infer_loop(self): recover_paused=self.control_state_machine.try_recover_paused_reqs(), ) + if self.support_overlap: + # The previous infer thread releases forward before its GPU work finishes. + # Keep the next DP collective behind that work to avoid overlapping NCCL + # with the previous DeepEP/MTP kernels. + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + dp_prefill_req_nums, dp_decode_req_nums = self._dp_all_gather_prefill_and_decode_req_num( prefill_reqs=prefill_reqs, decode_reqs=decode_reqs ) From 83b8dfa973ca1b1322a64b51d1d1e09e099ab2d3 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 11:12:06 +0800 Subject: [PATCH 137/214] feat: support DeepSeek-V4 DSpark --- lightllm/common/basemodel/basemodel.py | 1 + lightllm/common/basemodel/batch_objs.py | 4 + lightllm/common/basemodel/hidden_collector.py | 5 + lightllm/common/basemodel/infer_struct.py | 1 + .../deepseek4_mem_manager.py | 52 +++- lightllm/models/__init__.py | 1 + .../layer_infer/transformer_layer_infer.py | 1 + lightllm/models/deepseek_v4/model.py | 8 + .../triton_kernel/build_dspark_swa_index.py | 123 +++++++++ lightllm/models/deepseek_v4/workspace.py | 34 ++- .../models/deepseek_v4_dspark/__init__.py | 4 + .../models/deepseek_v4_dspark/infer_struct.py | 44 ++++ .../layer_infer/__init__.py | 16 ++ .../layer_infer/post_layer_infer.py | 90 +++++++ .../layer_infer/pre_layer_infer.py | 30 +++ .../layer_infer/transformer_layer_infer.py | 41 +++ .../layer_weights/__init__.py | 12 + .../pre_and_post_layer_weight.py | 130 +++++++++ .../layer_weights/transformer_layer_weight.py | 13 + lightllm/models/deepseek_v4_dspark/model.py | 249 ++++++++++++++++++ .../mtp_speculative/proposers/base.py | 7 + .../mtp_speculative/proposers/dspark.py | 15 +- .../model_infer/mtp_speculative/utils.py | 16 +- lightllm/utils/envs_utils.py | 19 ++ test/unit/test_deepseek_v4_dspark.py | 155 +++++++++++ 25 files changed, 1060 insertions(+), 11 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py create mode 100644 lightllm/models/deepseek_v4_dspark/__init__.py create mode 100644 lightllm/models/deepseek_v4_dspark/infer_struct.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_infer/__init__.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_infer/pre_layer_infer.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_weights/__init__.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py create mode 100644 lightllm/models/deepseek_v4_dspark/layer_weights/transformer_layer_weight.py create mode 100644 lightllm/models/deepseek_v4_dspark/model.py create mode 100644 test/unit/test_deepseek_v4_dspark.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 387864efb6..2ff8694618 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -430,6 +430,7 @@ def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0) # 特殊模型,特殊模式的特定变量初始化操作。 infer_state.mtp_draft_input_hiddens = model_input.mtp_draft_input_hiddens + infer_state.mtp_draft_swa_pages = model_input.mtp_draft_swa_pages if infer_state.is_prefill: infer_state.prefill_att_state = self.prefill_att_backend.create_att_prefill_state(infer_state=infer_state) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index c761ea4835..7d2b296253 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -57,6 +57,10 @@ class ModelInput: # mtp_draft_input_hiddens 用于模型 mtp 模式下 # 的 draft 模型的输入 mtp_draft_input_hiddens: Optional[torch.Tensor] = None + # DSpark draft block 的临时 SWA page 所有权。CPU tensor 用于无 D2H + # 回收;GPU tensor 供 attention 直接计算 block 的物理 SWA 槽。 + mtp_draft_swa_pages_cpu: Optional[torch.Tensor] = None + mtp_draft_swa_pages: Optional[torch.Tensor] = None # 主模型为 None: 准备所有 MTP 列;draft 首轮为 (): 无新槽;draft 追加后为 (k,): 只准备新槽。 mtp_decode_slot_prepare_indices: Optional[tuple] = None diff --git a/lightllm/common/basemodel/hidden_collector.py b/lightllm/common/basemodel/hidden_collector.py index 3eb946fe82..c8cedc19ff 100644 --- a/lightllm/common/basemodel/hidden_collector.py +++ b/lightllm/common/basemodel/hidden_collector.py @@ -206,6 +206,8 @@ def _load_layer_ids(self) -> frozenset[int]: layer_ids = draft_config.get("target_layer_ids") if layer_ids is None: layer_ids = draft_config.get("dflash_config", {}).get("target_layer_ids") + if layer_ids is None: + layer_ids = draft_config.get("dspark_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) @@ -217,6 +219,9 @@ def _load_layer_ids(self) -> frozenset[int]: def add(self, layer_index: int, hidden: torch.Tensor) -> None: if layer_index not in self.layer_ids: return + prepare_hidden = getattr(self.model, "prepare_mtp_layer_hidden", None) + if prepare_hidden is not None: + hidden = prepare_hidden(layer_index=layer_index, hidden=hidden) # 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()) diff --git a/lightllm/common/basemodel/infer_struct.py b/lightllm/common/basemodel/infer_struct.py index 91c6e99699..0df02e17db 100755 --- a/lightllm/common/basemodel/infer_struct.py +++ b/lightllm/common/basemodel/infer_struct.py @@ -94,6 +94,7 @@ def __init__(self): # 在开启 mtp_mode 时,mtp draft model # 的输入会用到,其他模型和场景都不会用到 self.mtp_draft_input_hiddens: Optional[torch.Tensor] = None + self.mtp_draft_swa_pages: Optional[torch.Tensor] = None # 在单节点多dp的运行模式下,在进行prefill的阶段,如果出现了dp之间数据不平衡的现象, # 可以将推理的数据,进行重新分配到各个dp,在做 att 之前,重新 all to all 到各自的 diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 305457d166..c2c5b787f3 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -901,6 +901,48 @@ def alloc_swa_decode( self._update_swa_page_counts(slots, 1) return + def alloc_dspark_swa_block(self, mem_indexes: torch.Tensor, block_size: int): + """Assign one temporary SWA scratch page to each DSpark proposal block. + + Proposal KV is consumed by the draft model immediately and every full + slot is released after the following target verify. Keeping the whole + block in a private page avoids the host-side sequence-length decision + required by the position-aligned target cache. Attention still uses + the absolute positions stored in the request table; physical SWA slots + only identify the packed KV rows. + """ + mem_indexes = mem_indexes.reshape(-1) + block_size = int(block_size) + assert block_size > 0 + assert block_size <= DSV4_SWA_PAGE_SIZE + assert mem_indexes.numel() % block_size == 0 + + req_num = mem_indexes.numel() // block_size + if req_num == 0: + return ( + torch.empty((0,), dtype=torch.int32, device="cpu"), + torch.empty((0,), dtype=torch.int32, device=self.full_to_swa_indexs.device), + ) + + device = self.full_to_swa_indexs.device + pages_cpu = self._alloc_swa_pages(req_num) + pages = pages_cpu.to(device, non_blocking=True) + # DSpark block 的物理槽由 attention index builder 直接从 page id + # 计算,不发布到 target 的全局 full->SWA 映射,也不参与 live count。 + return pages_cpu, pages + + def free_dspark_swa_block( + self, + mem_indexes_cpu: torch.Tensor, + pages_cpu: torch.Tensor, + ) -> None: + """Return DSpark scratch resources using their original CPU handles.""" + assert not mem_indexes_cpu.is_cuda + assert not pages_cpu.is_cuda + self.swa_page_allocator.free(pages_cpu) + MemoryManager.free(self, mem_indexes_cpu) + return + def evict_swa(self, full_slots: torch.Tensor) -> None: """回收 full 槽位对应的 swa 槽(出窗惰性回收 / free 级联 / 压力阀共用)。 未映射(-1)的槽位跳过;页计数减到 0 时整页归还 allocator。""" @@ -1107,6 +1149,7 @@ def pack_mla_kv_to_cache_fused_norm_rope( eps: float, freqs_cis: torch.Tensor, positions: torch.Tensor, + swa_slots: Optional[torch.Tensor] = None, ): """同 pack_mla_kv_to_cache,但 rmsnorm + 尾部交错 rope 融合进写入 kernel 并省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" @@ -1114,8 +1157,13 @@ def pack_mla_kv_to_cache_fused_norm_rope( fused_k_norm_rope_flashmla, ) - swa_slots = self.full_to_swa_indexs[mem_index.reshape(-1)] - swa_slots = torch.where(swa_slots < 0, torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), swa_slots) + if swa_slots is None: + swa_slots = self.full_to_swa_indexs[mem_index.reshape(-1)] + swa_slots = torch.where( + swa_slots < 0, + torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), + swa_slots, + ) fused_k_norm_rope_flashmla( kv=kv, kv_weight=kv_weight, diff --git a/lightllm/models/__init__.py b/lightllm/models/__init__.py index fd3a91c805..0d2d8f35de 100644 --- a/lightllm/models/__init__.py +++ b/lightllm/models/__init__.py @@ -46,6 +46,7 @@ from lightllm.models.qwen3_5_moe.model import Qwen3_5MOETpPartModel from lightllm.models.deepseek_mtp.model import Deepseek3MTPModel from lightllm.models.deepseek_v4_mtp.model import DeepseekV4MTPModel +from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel 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 diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index f31a65953b..ecb147e5a2 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -358,6 +358,7 @@ def _get_qkv( infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, mem_index=infer_state.mem_index, + swa_slots=getattr(infer_state, "dsv4_swa_write_slots", None), kv=qkv[:, -self.head_dim_ :], kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 0a152c384a..24f088c819 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -223,6 +223,14 @@ def _kernel_warmup(self): ) return + def prepare_mtp_layer_hidden(self, layer_index: int, hidden): + """Materialize the official per-layer DSpark feature from the deferred mHC state.""" + if isinstance(hidden, tuple): + streams = hc_post(*hidden) + else: + streams = hidden.view(-1, self.config["hc_mult"], self.config["hidden_size"]) + return streams.mean(dim=1) + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: """Commit DSV4 derived slots before BaseModel pads or scatters the generic input.""" if model_input.batch_size == 0: diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py new file mode 100644 index 0000000000..a320a15068 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py @@ -0,0 +1,123 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _build_dspark_swa_index_kernel( + req_idx_ptr, + pos_ptr, + req_to_token_ptr, + req_to_token_stride0, + full_to_swa_ptr, + scratch_pages_ptr, + swa_index_ptr, + swa_index_stride0, + swa_length_ptr, + swa_write_slot_ptr, + HOLD_REQ_ID: tl.constexpr, + HOLD_FULL_SLOT: tl.constexpr, + HOLD_SWA_SLOT: tl.constexpr, + WINDOW: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + PAGE_SIZE: tl.constexpr, + WIDTH: tl.constexpr, + BLOCK_W: tl.constexpr, +): + token_idx = tl.program_id(0) + req_idx = tl.load(req_idx_ptr + token_idx).to(tl.int64) + position = tl.load(pos_ptr + token_idx).to(tl.int64) + is_hold = req_idx == HOLD_REQ_ID + + block_offset = token_idx % BLOCK_SIZE + block_start = tl.maximum(position - block_offset, 0) + history_len = tl.minimum(block_start, WINDOW) + + column = tl.arange(0, BLOCK_W) + in_width = column < WIDTH + is_history = column < history_len + is_block = (column >= history_len) & (column < history_len + BLOCK_SIZE) + history_position = block_start - 1 - column + block_position = block_start + column - history_len + source_position = tl.where(is_history, history_position, block_position) + source_position = tl.where(is_hold | ~(is_history | is_block), 0, source_position) + + history_full_slot = tl.load( + req_to_token_ptr + req_idx * req_to_token_stride0 + source_position, + mask=in_width & is_history & ~is_hold, + other=HOLD_FULL_SLOT, + ).to(tl.int64) + history_swa_slot = tl.load( + full_to_swa_ptr + history_full_slot, + mask=in_width & is_history & ~is_hold, + other=HOLD_SWA_SLOT, + ) + scratch_page = tl.load( + scratch_pages_ptr + token_idx // BLOCK_SIZE, + mask=~is_hold, + other=0, + ).to(tl.int64) + block_swa_slot = scratch_page * PAGE_SIZE + column - history_len + swa_slot = tl.where(is_history, history_swa_slot, block_swa_slot) + valid = is_history | is_block + output = tl.where(valid, swa_slot, -1) + output = tl.where(is_hold, HOLD_SWA_SLOT, output).to(tl.int32) + tl.store( + swa_index_ptr + token_idx * swa_index_stride0 + column, output, mask=in_width + ) + + length = tl.where(is_hold, 1, history_len + BLOCK_SIZE).to(tl.int32) + tl.store(swa_length_ptr + token_idx, length) + write_slot = scratch_page * PAGE_SIZE + block_offset + write_slot = tl.where(is_hold, HOLD_SWA_SLOT, write_slot).to(tl.int32) + tl.store(swa_write_slot_ptr + token_idx, write_slot) + + +def build_dspark_swa_index( + req_idx: torch.Tensor, + positions: torch.Tensor, + req_to_token_indexs: torch.Tensor, + full_to_swa_indexs: torch.Tensor, + scratch_pages: torch.Tensor, + swa_index: torch.Tensor, + swa_length: torch.Tensor, + swa_write_slots: torch.Tensor, + window: int, + block_size: int, + page_size: int, + hold_req_id: int, + hold_full_slot: int, + hold_swa_slot: int, +): + """Build ``history SWA + complete draft block`` indices for every DSpark query row.""" + + token_num = positions.shape[0] + width = swa_index.shape[1] + assert token_num % block_size == 0 + assert scratch_pages is not None and scratch_pages.is_cuda + assert window > 0 and width >= window + block_size + if token_num == 0: + return swa_index, swa_length + + _build_dspark_swa_index_kernel[(token_num,)]( + req_idx, + positions, + req_to_token_indexs, + req_to_token_indexs.stride(0), + full_to_swa_indexs, + scratch_pages, + swa_index, + swa_index.stride(0), + swa_length, + swa_write_slots, + HOLD_REQ_ID=hold_req_id, + HOLD_FULL_SLOT=hold_full_slot, + HOLD_SWA_SLOT=hold_swa_slot, + WINDOW=window, + BLOCK_SIZE=block_size, + PAGE_SIZE=page_size, + WIDTH=width, + BLOCK_W=triton.next_power_of_2(width), + num_warps=4, + ) + return swa_index, swa_length diff --git a/lightllm/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py index 2b260af69a..c5b581b139 100644 --- a/lightllm/models/deepseek_v4/workspace.py +++ b/lightllm/models/deepseek_v4/workspace.py @@ -6,14 +6,26 @@ class DeepseekV4Workspace: def __init__(self, model): self.token_capacity = int(model.batch_max_tokens) self.sliding_window = int(model.config["sliding_window"]) + args = get_env_start_args() + self.swa_capacity = self.sliding_window + if args.mtp_mode == "dspark": + dspark_width = self.sliding_window + int(args.mtp_step) + # FlashMLA sparse decode requires the physical top-k width to be block aligned. + # 128 covers both supported padded Q-head configurations; swa_lengths keeps + # the actual number of visible history + draft-block entries. + self.swa_capacity = ((dspark_width + 127) // 128) * 128 self.index_topk = int(model.config["index_topk"]) self.c128_cap = self.compress_cap(model.max_seq_length, 128) - args = get_env_start_args() overlap = args.enable_decode_microbatch_overlap or args.enable_prefill_microbatch_overlap self.microbatch_count = 1 + int(overlap) - self.swa_indices = self._alloc(self.sliding_window) + self.swa_indices = self._alloc(self.swa_capacity) self.swa_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") + self.dspark_swa_write_slots = torch.empty( + (self.microbatch_count, self.token_capacity), + dtype=torch.int32, + device="cuda", + ) self.c4_indices = self._alloc(self.index_topk) self.c4_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") self.c128_indices = self._alloc(self.c128_cap) @@ -46,12 +58,26 @@ def _alloc(self, width: int) -> torch.Tensor: def _view(buffer: torch.Tensor, token_num: int, width: int) -> torch.Tensor: return torch.as_strided(buffer, (token_num, width), (width, 1)) - def swa(self, microbatch_index: int, token_num: int): + def swa(self, microbatch_index: int, token_num: int, width: int = None): + width = self.sliding_window if width is None else int(width) + assert width <= self.swa_capacity, f"swa width {width} exceeds allocated {self.swa_capacity}" return ( - self._view(self.swa_indices[microbatch_index], token_num, self.sliding_window), + self._view(self.swa_indices[microbatch_index], token_num, width), self.swa_lengths[microbatch_index, :token_num], ) + def dspark_swa(self, microbatch_index: int, token_num: int): + indices, lengths = self.swa( + microbatch_index, + token_num, + width=self.swa_capacity, + ) + return ( + indices, + lengths, + self.dspark_swa_write_slots[microbatch_index, :token_num], + ) + def c4(self, microbatch_index: int, token_num: int, width: int): assert width <= self.index_topk, f"c4 width {width} exceeds allocated {self.index_topk}" return ( diff --git a/lightllm/models/deepseek_v4_dspark/__init__.py b/lightllm/models/deepseek_v4_dspark/__init__.py new file mode 100644 index 0000000000..957345c91e --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/__init__.py @@ -0,0 +1,4 @@ +from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + +__all__ = ["DeepseekV4DSparkModel"] diff --git a/lightllm/models/deepseek_v4_dspark/infer_struct.py b/lightllm/models/deepseek_v4_dspark/infer_struct.py new file mode 100644 index 0000000000..73cb48c220 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/infer_struct.py @@ -0,0 +1,44 @@ +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_SWA_PAGE_SIZE, +) +from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.deepseek_v4.triton_kernel.build_dspark_swa_index import ( + build_dspark_swa_index, +) + + +class DeepseekV4DSparkInferStateInfo(DeepseekV4InferStateInfo): + """DeepSeek-V4 metadata with non-causal visibility across one DSpark block.""" + + def init_some_extra_state(self, model): + super().init_some_extra_state(model) + if self.is_prefill or self.mtp_draft_swa_pages is None: + # Target-hidden commit passes use target-owned mappings. Proposal + # blocks, including CUDA Graph HOLD capture, always carry scratch + # pages and replace these base indices below. + return + + ( + self.dsv4_swa_indices, + self.dsv4_swa_lengths, + self.dsv4_swa_write_slots, + ) = model.dsv4_workspace.dspark_swa( + self.microbatch_index, + self.position_ids.numel(), + ) + build_dspark_swa_index( + req_idx=self.dsv4_sparse_req_idx, + positions=self.position_ids, + req_to_token_indexs=self.req_manager.req_to_token_indexs, + full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, + scratch_pages=self.mtp_draft_swa_pages, + swa_index=self.dsv4_swa_indices, + swa_length=self.dsv4_swa_lengths, + swa_write_slots=self.dsv4_swa_write_slots, + window=model.config["sliding_window"], + block_size=model.block_size, + page_size=DSV4_SWA_PAGE_SIZE, + hold_req_id=self.req_manager.HOLD_REQUEST_ID, + hold_full_slot=self.mem_manager.HOLD_TOKEN_MEMINDEX, + hold_swa_slot=self.mem_manager.swa_pool.HOLD_TOKEN_MEMINDEX, + ) diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/__init__.py b/lightllm/models/deepseek_v4_dspark/layer_infer/__init__.py new file mode 100644 index 0000000000..8279c79850 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/__init__.py @@ -0,0 +1,16 @@ +from lightllm.models.deepseek_v4_dspark.layer_infer.post_layer_infer import ( + DeepseekV4DSparkPostLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_infer.pre_layer_infer import ( + DeepseekV4DSparkPreLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_infer.transformer_layer_infer import ( + DeepseekV4DSparkTransformerLayerInfer, +) + + +__all__ = [ + "DeepseekV4DSparkPostLayerInfer", + "DeepseekV4DSparkPreLayerInfer", + "DeepseekV4DSparkTransformerLayerInfer", +] diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py new file mode 100644 index 0000000000..002ca26fdd --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py @@ -0,0 +1,90 @@ +import torch + +from lightllm.models.deepseek_v4.layer_infer.hyper_connection import hc_head, hc_post +from lightllm.models.deepseek_v4_dspark.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4DSparkPreAndPostLayerWeight, +) +from lightllm.models.qwen3_dspark.layer_infer.post_layer_infer import ( + Qwen3DSparkPostLayerInfer, +) + + +class DeepseekV4DSparkPostLayerInfer(Qwen3DSparkPostLayerInfer): + """Collapse mHC streams, then run the shared head plus DSpark Markov/confidence heads.""" + + @torch.no_grad() + def predict_confidence_logits( + self, + block_hidden: torch.Tensor, + anchor_token_ids: torch.Tensor, + sampled_tokens: torch.Tensor, + layer_weight: DeepseekV4DSparkPreAndPostLayerWeight, + ): + 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) + features = torch.cat( + [block_hidden, prev_embeddings.to(dtype=block_hidden.dtype)], dim=-1 + ) + logits = layer_weight.confidence_head_weight_.mm( + features.flatten(0, -2).float() + ) + return logits.view(features.shape[:-1]) + + def token_forward( + self, + input_embdings: torch.Tensor, + infer_state, + layer_weight: DeepseekV4DSparkPreAndPostLayerWeight, + ): + if infer_state.is_prefill: + return input_embdings.new_empty((0,)) + + if isinstance(input_embdings, tuple): + streams = hc_post(*input_embdings) + input_embdings = streams.reshape(streams.shape[0], -1) + collapsed = hc_head( + input_embdings, + layer_weight.hc_head_fn_.weight, + layer_weight.hc_head_scale_.weight, + layer_weight.hc_head_base_.weight, + layer_weight.network_config_["hc_mult"], + layer_weight.network_config_["hidden_size"], + layer_weight.network_config_["rms_norm_eps"], + layer_weight.network_config_.get("hc_eps", 1e-6), + self.alloc_tensor, + ) + + last_input, token_num = self._slice_get_last_input(collapsed, 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 + ] + + 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, + ) + # The proposer consumes token ids directly; keep only the row dimension for graph unpadding. + return local_logits.new_empty((token_num, 1)) diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/pre_layer_infer.py new file mode 100644 index 0000000000..4de7acad83 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/pre_layer_infer.py @@ -0,0 +1,30 @@ +from lightllm.models.deepseek_v4.layer_infer.pre_layer_infer import ( + DeepseekV4PreLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4DSparkPreAndPostLayerWeight, +) + + +class DeepseekV4DSparkPreLayerInfer(DeepseekV4PreLayerInfer): + """Project selected target-layer hiddens for draft KV commits.""" + + 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: DeepseekV4DSparkPreAndPostLayerWeight, + ): + target_hidden = layer_weight.main_proj_weight_.mm( + infer_state.mtp_draft_input_hiddens, + use_custom_tensor_mananger=False, + ) + return layer_weight.main_norm_weight_( + input=target_hidden, + eps=self.eps_, + alloc_func=self.alloc_tensor, + ) diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py new file mode 100644 index 0000000000..9006434671 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py @@ -0,0 +1,41 @@ +import torch + +from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import ( + DeepseekV4TransformerLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_weights.transformer_layer_weight import ( + DeepseekV4DSparkTransformerLayerWeight, +) + + +class DeepseekV4DSparkTransformerLayerInfer(DeepseekV4TransformerLayerInfer): + """Full DSpark stage with a KV-only target-hidden commit primitive.""" + + def __init__(self, layer_num, network_config): + super().__init__(layer_num, network_config) + final_layer = network_config["n_layer"] + network_config["dspark_layer_num"] - 1 + self.is_last_layer = layer_num == final_layer + assert ( + self.compress_ratio == 0 + ), "DeepSeek-V4 DSpark draft layers must be SWA-only" + + def context_forward( + self, + input_embdings: torch.Tensor, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4DSparkTransformerLayerWeight, + ) -> torch.Tensor: + """Write target hidden rows into this stage without running draft attention/FFN.""" + full_input = self._tpsp_allgather(input=input_embdings, infer_state=infer_state) + qkv = layer_weight.wq_a_wkv_.mm(full_input, use_custom_tensor_mananger=False) + infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( + layer_index=self.layer_num_, + mem_index=infer_state.mem_index, + kv=qkv[:, -self.head_dim_ :], + kv_weight=layer_weight.kv_norm_.weight, + eps=self.eps_, + freqs_cis=self.freqs_cis, + positions=infer_state.position_ids, + ) + return input_embdings diff --git a/lightllm/models/deepseek_v4_dspark/layer_weights/__init__.py b/lightllm/models/deepseek_v4_dspark/layer_weights/__init__.py new file mode 100644 index 0000000000..ce1b9f3b04 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_weights/__init__.py @@ -0,0 +1,12 @@ +from lightllm.models.deepseek_v4_dspark.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4DSparkPreAndPostLayerWeight, +) +from lightllm.models.deepseek_v4_dspark.layer_weights.transformer_layer_weight import ( + DeepseekV4DSparkTransformerLayerWeight, +) + + +__all__ = [ + "DeepseekV4DSparkPreAndPostLayerWeight", + "DeepseekV4DSparkTransformerLayerWeight", +] diff --git a/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py new file mode 100644 index 0000000000..f5159f563e --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py @@ -0,0 +1,130 @@ +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 +from lightllm.models.deepseek_v4.triton_kernel.quant_convert import ( + dequant_fp8_block_to_bf16, +) + + +class DeepseekV4DSparkPreAndPostLayerWeight(PreAndPostLayerWeight): + """Shared vocabulary weights plus the first/last DSpark stage heads.""" + + def __init__(self, data_type, network_config, quant_cfg: Quantcfg): + super().__init__(data_type, network_config) + self.quant_cfg = quant_cfg + + hidden = network_config["hidden_size"] + vocab = network_config["vocab_size"] + hc_mult = network_config["hc_mult"] + markov_rank = network_config["dspark_markov_rank"] + target_hidden_size = hidden * len(network_config["dspark_target_layer_ids"]) + first_layer_idx = network_config["n_layer"] + last_stage = network_config["dspark_layer_num"] - 1 + last_prefix = f"mtp.{last_stage}" + + # Replaced with the already-loaded target weights by the model. + self.wte_weight_ = EmbeddingWeight( + dim=hidden, + vocab_size=vocab, + weight_name="embed.weight", + data_type=self.data_type_, + ) + self.lm_head_weight_ = LMHeadWeight( + dim=hidden, + vocab_size=vocab, + weight_name="head.weight", + data_type=self.data_type_, + ) + + self.main_proj_weight_ = ROWMMWeight( + in_dim=target_hidden_size, + out_dims=[hidden], + weight_names="mtp.0.main_proj.weight", + data_type=self.data_type_, + quant_method=self.quant_cfg.get_quant_method(first_layer_idx, "main_proj"), + tp_rank=0, + tp_world_size=1, + ) + self.main_norm_weight_ = RMSNormWeight( + dim=hidden, + weight_name="mtp.0.main_norm.weight", + data_type=self.data_type_, + ) + + self.final_norm_weight_ = RMSNormWeight( + dim=hidden, + weight_name=f"{last_prefix}.norm.weight", + data_type=self.data_type_, + ) + self.hc_head_fn_ = ParameterWeight( + weight_name=f"{last_prefix}.hc_head_fn", + data_type=torch.float32, + weight_shape=(hc_mult, hc_mult * hidden), + ) + self.hc_head_base_ = ParameterWeight( + weight_name=f"{last_prefix}.hc_head_base", + data_type=torch.float32, + weight_shape=(hc_mult,), + ) + self.hc_head_scale_ = ParameterWeight( + weight_name=f"{last_prefix}.hc_head_scale", + data_type=torch.float32, + weight_shape=(1,), + ) + + # W1 is replicated to avoid one TP collective in every sequential Markov step. + self.markov_w1_weight_ = EmbeddingWeight( + dim=markov_rank, + vocab_size=vocab, + weight_name=f"{last_prefix}.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, + weight_name=f"{last_prefix}.markov_head.markov_w2.weight", + data_type=self.data_type_, + ) + self.confidence_head_weight_ = ROWMMWeight( + in_dim=hidden + markov_rank, + out_dims=[1], + weight_names=f"{last_prefix}.confidence_head.proj.weight", + bias_names=None, + data_type=torch.float32, + quant_method=None, + tp_rank=0, + tp_world_size=1, + ) + + # Reused by the generic DSpark Markov implementation. + self.markov_rank = markov_rank + self.markov_head_type = "vanilla" + + def load_hf_weights(self, weights): + self._dequant_main_proj_in_place(weights) + return super().load_hf_weights(weights) + + def _dequant_main_proj_in_place(self, weights): + weight_name = self.main_proj_weight_.weight_names[0] + scale_key = weight_name[: -len(".weight")] + ".scale" + if scale_key not in weights: + return + + scale_name = self.main_proj_weight_.weight_scale_names[0] + if scale_name is None: + weights[weight_name] = dequant_fp8_block_to_bf16( + weights[weight_name], weights[scale_key] + ).to(self.data_type_) + else: + weights[scale_name] = weights[scale_key].to(torch.float32) + del weights[scale_key] diff --git a/lightllm/models/deepseek_v4_dspark/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4_dspark/layer_weights/transformer_layer_weight.py new file mode 100644 index 0000000000..b92b5637b3 --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/layer_weights/transformer_layer_weight.py @@ -0,0 +1,13 @@ +from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import ( + DeepseekV4TransformerLayerWeight, +) + + +class DeepseekV4DSparkTransformerLayerWeight(DeepseekV4TransformerLayerWeight): + """A DSpark stage whose cache layer id is global but checkpoint prefix is local.""" + + def _parse_config(self): + super()._parse_config() + stage_id = self.layer_num_ - self.network_config_["n_layer"] + assert 0 <= stage_id < self.network_config_["dspark_layer_num"] + self.prefix = f"mtp.{stage_id}" diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py new file mode 100644 index 0000000000..aaf4765e1a --- /dev/null +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -0,0 +1,249 @@ +import gc +import os +from typing import List + +import torch +from safetensors import safe_open +from tqdm import tqdm + +import lightllm.utils.petrel_helper as utils +from lightllm.common.basemodel import TpPartBaseModel +from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel +from lightllm.models.deepseek_v4_dspark.infer_struct import ( + DeepseekV4DSparkInferStateInfo, +) +from lightllm.models.deepseek_v4_dspark.layer_infer.post_layer_infer import ( + DeepseekV4DSparkPostLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_infer.pre_layer_infer import ( + DeepseekV4DSparkPreLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_infer.transformer_layer_infer import ( + DeepseekV4DSparkTransformerLayerInfer, +) +from lightllm.models.deepseek_v4_dspark.layer_weights.pre_and_post_layer_weight import ( + DeepseekV4DSparkPreAndPostLayerWeight, +) +from lightllm.models.deepseek_v4_dspark.layer_weights.transformer_layer_weight import ( + DeepseekV4DSparkTransformerLayerWeight, +) +from lightllm.models.draft_registry import DraftModelRegistry +from lightllm.utils.log_utils import init_logger + + +logger = init_logger(__name__) + + +@DraftModelRegistry(model_type="deepseek_v4", spec_modes="dspark") +class DeepseekV4DSparkModel(DeepseekV4TpPartModel): + """Three-stage DSpark draft model backed by target-owned DeepSeek-V4 cache layers.""" + + is_mtp_draft_model = True + + pre_and_post_weight_class = DeepseekV4DSparkPreAndPostLayerWeight + transformer_weight_class = DeepseekV4DSparkTransformerLayerWeight + pre_layer_infer_class = DeepseekV4DSparkPreLayerInfer + post_layer_infer_class = DeepseekV4DSparkPostLayerInfer + transformer_layer_infer_class = DeepseekV4DSparkTransformerLayerInfer + infer_state_class = DeepseekV4DSparkInferStateInfo + + 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" + ) + assert ( + not self.mtp_previous_draft_models + ), "DeepSeek-V4 DSpark uses one parallel draft model" + + def _init_config(self): + super()._init_config() + draft_layer_num = len(self.config["compress_ratios"]) - self.config["n_layer"] + assert draft_layer_num > 0 + assert ( + self.config["compress_ratios"][-draft_layer_num:] == [0] * draft_layer_num + ) + assert len(self.config["dspark_target_layer_ids"]) == draft_layer_num + self.config["dspark_layer_num"] = draft_layer_num + + # Aliases consumed by the shared block/Markov implementation. + self.config["block_size"] = int(self.config["dspark_block_size"]) + self.config["mask_token_id"] = int(self.config["dspark_noise_token_id"]) + self.config["markov_rank"] = int(self.config["dspark_markov_rank"]) + self.config["markov_head_type"] = "vanilla" + self.config["enable_confidence_head"] = True + self.config["confidence_head_with_markov"] = True + + def _verify_params(self): + super()._verify_params() + assert ( + not self.enable_tpsp_mix_mode + ), "DeepSeek-V4 DSpark draft model does not support TP-SP" + assert self.args.mtp_step == self.config["dspark_block_size"], ( + f"DeepSeek-V4 DSpark requires --mtp_step {self.config['dspark_block_size']}, " + f"got {self.args.mtp_step}" + ) + + def _init_quant(self): + super()._init_quant() + expert_quant_type = self.quant_cfg.get_quant_type( + self.config["n_layer"] - 1, "fused_moe" + ) + for layer_idx in range( + self.config["n_layer"], len(self.config["compress_ratios"]) + ): + self.quant_cfg.quant_cfg[layer_idx]["fused_moe"] = expert_quant_type + + def _init_weights(self, start_layer_index=None): + assert start_layer_index is None + self.pre_post_weight = self.pre_and_post_weight_class( + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + self.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_ + ) + first_layer = self.config["n_layer"] + self.trans_layers_weight = [ + self.transformer_weight_class( + first_layer + stage_id, + self.data_type, + network_config=self.config, + quant_cfg=self.quant_cfg, + ) + for stage_id in range(self.config["dspark_layer_num"]) + ] + + 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_infer_layer(self, start_layer_index=None): + assert start_layer_index is None + self.pre_infer = self.pre_layer_infer_class(network_config=self.config) + self.post_infer = self.post_layer_infer_class(network_config=self.config) + first_layer = self.config["n_layer"] + self.layers_infer = [ + self.transformer_layer_infer_class( + first_layer + stage_id, network_config=self.config + ) + for stage_id in range(self.config["dspark_layer_num"]) + ] + + def _init_some_value(self): + super()._init_some_value() + self.layers_num = self.config["dspark_layer_num"] + + def _init_custom(self): + self._freqs_cis_sliding = self.main_model._freqs_cis_sliding + self._freqs_cis_compress = self.main_model._freqs_cis_compress + self._cos_cached_sliding = self.main_model._cos_cached_sliding + self._sin_cached_sliding = self.main_model._sin_cached_sliding + self._cos_cached_compress = self.main_model._cos_cached_compress + self._sin_cached_compress = self.main_model._sin_cached_compress + self.dsv4_workspace = self.main_model.dsv4_workspace + self.block_size = self.config["dspark_block_size"] + self.mask_token_id = self.config["dspark_noise_token_id"] + for layer in self.layers_infer: + layer.freqs_cis = self._freqs_cis_sliding + layer.cos_compress_table = self._cos_cached_compress + layer.sin_compress_table = self._sin_cached_compress + + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: + if not model_input.is_prefill: + if model_input.mtp_draft_input_hiddens is not None: + # Target verify already prepared these shared full/SWA slots. + # This pass only commits target hiddens into the draft layers. + model_input.mtp_decode_slot_prepare_indices = () + elif model_input.mem_indexes_cpu is not None: + # Proposal-owned full slots live only until the next verify. + # Put each request's complete block in a private scratch page; + # this needs neither accepted-length D2H nor host seq metadata. + ( + model_input.mtp_draft_swa_pages_cpu, + model_input.mtp_draft_swa_pages, + ) = self.mem_manager.alloc_dspark_swa_block( + mem_indexes=model_input.mem_indexes, + block_size=self.block_size, + ) + model_input.mtp_decode_slot_prepare_indices = () + return super()._prepare_dsv4_slots(model_input) + + 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 + infer_state = self._create_inferstate(model_input) + infer_state.position_ids = model_input.b_seq_len - 1 + + 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): + assert token_num % self.block_size == 0 + return { + "mtp_draft_input_hiddens": None, + # CUDA Graph capture must take the same direct scratch-slot branch + # as runtime. HOLD rows never dereference the page value, but the + # tensor gives the captured graph stable storage for replay copies. + "mtp_draft_swa_pages": torch.zeros( + token_num // self.block_size, dtype=torch.int32, device="cuda" + ), + } + + def _autotune_warmup(self): + return + + def _init_padded_req(self): + return + + def _init_prefill_cuda_graph(self): + self.prefill_graph = None + + def load_weights(self, weight_dict: dict): + if weight_dict: + return super().load_weights(weight_dict) + + index_file = os.path.join(self.weight_dir_, "model.safetensors.index.json") + assert utils.PetrelHelper.exists( + index_file + ), "DeepSeek-V4 DSpark requires model.safetensors.index.json" + weight_map = utils.PetrelHelper.load_json(index_file)["weight_map"] + candidate_files = sorted( + {file_ for key, file_ in weight_map.items() if key.startswith("mtp.")} + ) + assert ( + candidate_files + ), "DeepSeek-V4 DSpark weights with prefix mtp.* were not found" + + loaded_key_count = 0 + desc = f"pid {os.getpid()} Loading DeepSeek-V4 DSpark weights" + for file_ in tqdm(candidate_files, total=len(candidate_files), desc=desc): + weights = {} + with safe_open(os.path.join(self.weight_dir_, file_), "pt", "cpu") as f: + for key in f.keys(): + if key.startswith("mtp."): + weights[key] = f.get_tensor(key) + + loaded_key_count += len(weights) + self.pre_post_weight.load_hf_weights(weights) + for layer in self.trans_layers_weight: + layer.load_hf_weights(weights) + del weights + gc.collect() + + self.pre_post_weight.verify_load() + [weight.verify_load() for weight in self.trans_layers_weight] + logger.info("loaded DeepSeek-V4 DSpark weights: %d tensors", loaded_key_count) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py index 4b3d0ed6c6..ecb792ef47 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/base.py @@ -21,8 +21,15 @@ class MtpMemIndexesToFree: # 与 mem_indexes_cpu 形状一致的 bool Tensor;True 表示释放对应索引。 # 为 None 时表示 mem_indexes_cpu 中的全部索引都需要释放。 free_mask_cpu: Optional[torch.Tensor] = None + # DSpark proposal 独占的 SWA scratch pages。page ids 从 CPU allocator + # 产生并一直保留在 CPU;非 None 时走直接回收,不再从 GPU 映射反查。 + swa_pages_cpu: Optional[torch.Tensor] = None def __post_init__(self) -> None: + if self.swa_pages_cpu is not None: + assert isinstance(self.swa_pages_cpu, torch.Tensor) + assert not self.swa_pages_cpu.is_cuda + assert self.swa_pages_cpu.dtype == torch.int32 if self.free_mask_cpu is None: return assert isinstance(self.free_mask_cpu, torch.Tensor) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 5e3ea5694f..3144006f60 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -57,6 +57,8 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, + accept_len_cpu: torch.Tensor | None = None, + accept_len_ready_event: torch.cuda.Event | None = None, ) -> DSparkSpecProposal: """提交 target verify KV,并生成下一轮 DSpark block proposal。 @@ -144,8 +146,12 @@ def propose_next( .contiguous() ) draft_input.mem_indexes = extra_mem_indexes_cpu.cuda(non_blocking=True) - draft_input.mem_indexes_cpu = None + draft_input.mem_indexes_cpu = extra_mem_indexes_cpu + draft_input.mtp_draft_swa_pages_cpu = None + draft_input.mtp_draft_swa_pages = None + draft_input.mtp_decode_slot_prepare_indices = 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: @@ -184,7 +190,12 @@ def propose_next( return DSparkSpecProposal( token_ids=proposal_token_ids, - extra_mem_indexes_cpu=[MtpMemIndexesToFree(mem_indexes_cpu=extra_mem_indexes_cpu)], + extra_mem_indexes_cpu=[ + MtpMemIndexesToFree( + mem_indexes_cpu=extra_mem_indexes_cpu, + swa_pages_cpu=draft_input.mtp_draft_swa_pages_cpu, + ) + ], schedule_scores=schedule_scores, schedule_scores_cpu=schedule_scores_cpu, ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index 935fe92d33..de4758ebee 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -119,15 +119,25 @@ def free_mem_indexes( """Free all KV indexes described by the unified MTP memory list.""" mem_indexes_to_free = [] + dspark_scratch_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 extra_mem_to_free.swa_pages_cpu is None: + mem_indexes_to_free.append(extra_indexes_cpu) + else: + assert extra_mem_to_free.free_mask_cpu is None + dspark_scratch_to_free.append( + (extra_indexes_cpu, extra_mem_to_free.swa_pages_cpu) + ) + + mem_manager = backend.model.req_manager.mem_manager + for mem_indexes_cpu, pages_cpu in dspark_scratch_to_free: + mem_manager.free_dspark_swa_block(mem_indexes_cpu, pages_cpu) if mem_indexes_to_free: - backend.model.req_manager.mem_manager.free(torch.cat(mem_indexes_to_free, dim=0)) + mem_manager.free(torch.cat(mem_indexes_to_free, dim=0)) __all__ = [ diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 4bdd359b68..bdee6bcb2a 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -340,6 +340,25 @@ def get_mtp_weight_layer_num() -> int: 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) + + if draft_config.get("model_type") == "deepseek_v4" and draft_config.get("dspark_block_size"): + target_layer_num = draft_config.get("num_hidden_layers", draft_config.get("n_layer")) + compress_ratios = draft_config.get("compress_ratios") + if target_layer_num is not None and isinstance(compress_ratios, list): + draft_layer_num = len(compress_ratios) - int(target_layer_num) + if draft_layer_num > 0: + draft_ratios = compress_ratios[-draft_layer_num:] + assert all( + int(ratio) == 0 for ratio in draft_ratios + ), f"DeepSeek-V4 DSpark draft layers must be SWA-only, got {draft_ratios}" + target_layer_ids = draft_config.get("dspark_target_layer_ids") + if target_layer_ids is not None: + assert len(target_layer_ids) == draft_layer_num, ( + f"DeepSeek-V4 DSpark target layer count {len(target_layer_ids)} does not match " + f"draft layer count {draft_layer_num}" + ) + return draft_layer_num + # 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 diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py new file mode 100644 index 0000000000..4d09fc325b --- /dev/null +++ b/test/unit/test_deepseek_v4_dspark.py @@ -0,0 +1,155 @@ +import json +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.utils.envs_utils import _get_mtp_draft_backbone_layer_num + + +def test_deepseek_v4_dspark_layer_count_uses_trailing_swa_layers(tmp_path): + config = { + "model_type": "deepseek_v4", + "num_hidden_layers": 43, + "compress_ratios": [4] * 43 + [0, 0, 0], + "dspark_block_size": 5, + "dspark_target_layer_ids": [40, 41, 42], + } + (tmp_path / "config.json").write_text(json.dumps(config)) + + assert _get_mtp_draft_backbone_layer_num(str(tmp_path)) == 3 + + +def test_dspark_scratch_cleanup_keeps_page_ids_on_cpu(): + 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, + ) + + class FakeMemManager: + def __init__(self): + self.scratch_frees = [] + self.normal_frees = [] + + def free_dspark_swa_block(self, mem_indexes_cpu, pages_cpu): + self.scratch_frees.append((mem_indexes_cpu.clone(), pages_cpu.clone())) + + def free(self, mem_indexes_cpu): + self.normal_frees.append(mem_indexes_cpu.clone()) + + mem_manager = FakeMemManager() + backend = SimpleNamespace(model=SimpleNamespace(req_manager=SimpleNamespace(mem_manager=mem_manager))) + scratch_full = torch.tensor([10, 11, 12, 13, 14], dtype=torch.int32) + scratch_pages = torch.tensor([7], dtype=torch.int32) + normal_full = torch.tensor([20, 21], dtype=torch.int32) + + mtp_utils.free_mem_indexes( + backend=backend, + extra_mem_indexes_cpu=[ + MtpMemIndexesToFree( + mem_indexes_cpu=scratch_full, + swa_pages_cpu=scratch_pages, + ), + MtpMemIndexesToFree(mem_indexes_cpu=normal_full), + ], + ) + + assert len(mem_manager.scratch_frees) == 1 + torch.testing.assert_close(mem_manager.scratch_frees[0][0], scratch_full) + torch.testing.assert_close(mem_manager.scratch_frees[0][1], scratch_pages) + assert len(mem_manager.normal_frees) == 1 + torch.testing.assert_close(mem_manager.normal_frees[0], normal_full) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_build_dspark_swa_index_exposes_history_and_complete_block(): + from lightllm.models.deepseek_v4.triton_kernel.build_dspark_swa_index import ( + build_dspark_swa_index, + ) + + block_size = 3 + window = 4 + req_to_token = torch.tensor( + [ + list(range(0, 10)), + list(range(10, 20)), + [20] * 10, + ], + dtype=torch.int32, + device="cuda", + ) + full_to_swa = torch.arange(21, dtype=torch.int32, device="cuda") + 100 + req_idx = torch.tensor([0, 0, 0, 1, 1, 1], dtype=torch.int32, device="cuda") + positions = torch.tensor([4, 5, 6, 2, 3, 4], dtype=torch.int32, device="cuda") + padded_width = 8 + indices = torch.empty((6, padded_width), dtype=torch.int32, device="cuda") + lengths = torch.empty((6,), dtype=torch.int32, device="cuda") + write_slots = torch.empty((6,), dtype=torch.int32, device="cuda") + scratch_pages = torch.tensor([2, 3], dtype=torch.int32, device="cuda") + + build_dspark_swa_index( + req_idx=req_idx, + positions=positions, + req_to_token_indexs=req_to_token, + full_to_swa_indexs=full_to_swa, + scratch_pages=scratch_pages, + swa_index=indices, + swa_length=lengths, + swa_write_slots=write_slots, + window=window, + block_size=block_size, + page_size=128, + hold_req_id=2, + hold_full_slot=20, + hold_swa_slot=120, + ) + + expected_indices = torch.tensor( + [ + [103, 102, 101, 100, 256, 257, 258, -1], + [103, 102, 101, 100, 256, 257, 258, -1], + [103, 102, 101, 100, 256, 257, 258, -1], + [111, 110, 384, 385, 386, -1, -1, -1], + [111, 110, 384, 385, 386, -1, -1, -1], + [111, 110, 384, 385, 386, -1, -1, -1], + ], + dtype=torch.int32, + device="cuda", + ) + torch.testing.assert_close(indices, expected_indices) + torch.testing.assert_close( + lengths, torch.tensor([7, 7, 7, 5, 5, 5], dtype=torch.int32, device="cuda") + ) + torch.testing.assert_close( + write_slots, + torch.tensor([256, 257, 258, 384, 385, 386], dtype=torch.int32, device="cuda"), + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_dspark_swa_block_uses_one_scratch_page_per_request(): + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DeepseekV4MemoryManager, + ) + + manager = DeepseekV4MemoryManager.__new__(DeepseekV4MemoryManager) + manager.full_to_swa_indexs = torch.full((32,), -1, dtype=torch.int32, device="cuda") + manager.swa_page_live_count = torch.zeros((4,), dtype=torch.int32, device="cuda") + manager._alloc_swa_pages = lambda count: torch.tensor( + [2, 0], + dtype=torch.int32, + pin_memory=True, + ) + mem_indexes = torch.tensor([3, 4, 5, 8, 9, 10], dtype=torch.int64, device="cuda") + + pages_cpu, pages = manager.alloc_dspark_swa_block(mem_indexes=mem_indexes, block_size=3) + + torch.testing.assert_close( + manager.full_to_swa_indexs[mem_indexes], + torch.full((6,), -1, dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + manager.swa_page_live_count, + torch.zeros((4,), dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close(pages_cpu, torch.tensor([2, 0], dtype=torch.int32)) From 16ca39bd036a8bed39031a8c56b85fd2e8600e4b Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 12:39:57 +0800 Subject: [PATCH 138/214] feat: add DeepSeek-V4 MTP BF16 conversion support --- lightllm/models/deepseek_v4/model.py | 8 +- lightllm/utils/config_utils.py | 21 + lightllm/utils/kv_cache_utils.py | 3 +- test/unit/test_deepseek_v4_dspark.py | 14 + tools/convert_deepseek_v4_mtp_to_bf16.py | 690 +++++++++++++++++++++++ 5 files changed, 732 insertions(+), 4 deletions(-) create mode 100644 tools/convert_deepseek_v4_mtp_to_bf16.py diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 24f088c819..5f236d2e4b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -40,7 +40,10 @@ linear_ramp_mask, ) from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num, get_env_start_args -from lightllm.utils.config_utils import normalize_deepseek_v4_config +from lightllm.utils.config_utils import ( + get_deepseek_v4_compress_rates, + normalize_deepseek_v4_config, +) from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager @@ -88,8 +91,7 @@ def _init_req_manager(self): return def _get_compress_rates(self, layer_num): - rates = list(self.config["compress_ratios"]) - return rates[:layer_num] + return get_deepseek_v4_compress_rates(self.config, layer_num) def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 7b2da8308a..16586f5b8c 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -54,6 +54,27 @@ def normalize_deepseek_v4_config(config: Dict[str, Any]) -> Dict[str, Any]: return config +def get_deepseek_v4_compress_rates(config: Dict[str, Any], layer_num: int) -> List[int]: + """Build KV compression rates for target and runtime-added draft layers.""" + target_layer_num = config.get("n_layer", config.get("num_hidden_layers")) + if target_layer_num is None: + raise ValueError("DeepSeek-V4 config is missing num_hidden_layers") + target_layer_num = int(target_layer_num) + rates = list(config["compress_ratios"]) + if len(rates) < target_layer_num: + raise ValueError( + f"DeepSeek-V4 compress_ratios has {len(rates)} entries, " + f"but the target model has {target_layer_num} layers" + ) + + if len(rates) < layer_num: + # The draft layer count comes from the selected draft model and need not match + # num_nextn_predict_layers in the independently trained target. Missing draft + # KV layers use sliding-window attention, whose compression ratio is zero. + rates.extend([0] * (layer_num - len(rates))) + return rates[:layer_num] + + def get_config_json(model_path: str): with open(os.path.join(model_path, "config.json"), "r") as file: json_obj = json.load(file) diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index 8fc1714be6..127812dc24 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -17,6 +17,7 @@ ) from lightllm.utils.log_utils import init_logger from lightllm.utils.config_utils import ( + get_deepseek_v4_compress_rates, get_config_json, get_num_key_value_heads, get_head_dim, @@ -128,7 +129,7 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": config = get_config_json(args.model_dir) layer_num = get_layer_num(args.model_dir) + get_added_mtp_kv_layer_num() layout = DeepseekV4CpuCacheLayout.from_compress_rates( - compress_rates=config["compress_ratios"][:layer_num], + compress_rates=get_deepseek_v4_compress_rates(config, layer_num), token_page_size=args.cpu_cache_token_page_size, head_dim=get_head_dim(args.model_dir), indexer_head_dim=config["index_head_dim"], diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index 4d09fc325b..ec06ac9dec 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -5,6 +5,7 @@ import torch from lightllm.utils.envs_utils import _get_mtp_draft_backbone_layer_num +from lightllm.utils.config_utils import get_deepseek_v4_compress_rates def test_deepseek_v4_dspark_layer_count_uses_trailing_swa_layers(tmp_path): @@ -20,6 +21,19 @@ def test_deepseek_v4_dspark_layer_count_uses_trailing_swa_layers(tmp_path): assert _get_mtp_draft_backbone_layer_num(str(tmp_path)) == 3 +def test_deepseek_v4_dspark_extends_target_compress_rates_for_draft_layers(): + target_rates = [0, 0] + [value for _ in range(20) for value in (4, 128)] + [4] + config = { + "num_hidden_layers": 43, + "num_nextn_predict_layers": 1, + "compress_ratios": target_rates + [0], + } + + rates = get_deepseek_v4_compress_rates(config, layer_num=46) + + assert rates == target_rates + [0, 0, 0] + + def test_dspark_scratch_cleanup_keeps_page_ids_on_cpu(): from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( diff --git a/tools/convert_deepseek_v4_mtp_to_bf16.py b/tools/convert_deepseek_v4_mtp_to_bf16.py new file mode 100644 index 0000000000..d0a2cfcb5b --- /dev/null +++ b/tools/convert_deepseek_v4_mtp_to_bf16.py @@ -0,0 +1,690 @@ +#!/usr/bin/env python3 +"""Convert DeepSeek-V4 ``mtp.*`` checkpoint tensors to BF16. + +The DeepSeek-V4-Flash checkpoint stores dense MTP matrices as block-FP8 and +routed-expert matrices as packed MXFP4. Casting those tensors directly would +not recover their values. This tool performs the corresponding dequantization, +removes the paired scale tensors from the output index, and writes every MTP +tensor as BF16. + +Non-MTP shards are linked or copied into a new model directory. The source +model is never modified. This changes checkpoint storage only; the runtime +must select the no-quantization path for the converted MTP layers while keeping +the target layers on their original quantization path. +""" + +from __future__ import annotations + +import argparse +import gc +import json +import math +import os +import shutil +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Iterable, List, Mapping, Sequence, Tuple + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + + +FP4_VALUES = ( + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + -0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, +) + +_DTYPE_BYTES = { + "BOOL": 1, + "U8": 1, + "I8": 1, + "F8_E4M3": 1, + "F8_E5M2": 1, + "F8_E8M0": 1, + "I16": 2, + "U16": 2, + "F16": 2, + "BF16": 2, + "I32": 4, + "U32": 4, + "F32": 4, + "I64": 8, + "U64": 8, + "F64": 8, +} + + +@dataclass(frozen=True) +class TensorSpec: + name: str + source_shard: str + source_dtype: str + source_shape: Tuple[int, ...] + output_shape: Tuple[int, ...] + output_nbytes: int + conversion: str + scale_name: str | None = None + + +def _numel(shape: Sequence[int]) -> int: + return math.prod(int(dim) for dim in shape) + + +def _tensor_nbytes(dtype: str, shape: Sequence[int]) -> int: + try: + item_size = _DTYPE_BYTES[dtype] + except KeyError as exc: + raise ValueError(f"unsupported safetensors dtype: {dtype}") from exc + return _numel(shape) * item_size + + +def _scale_name(weight_name: str) -> str: + if not weight_name.endswith(".weight"): + raise ValueError(f"quantized tensor does not end in .weight: {weight_name}") + return weight_name[: -len(".weight")] + ".scale" + + +def _parse_size(value: str) -> int: + text = value.strip().upper().replace(" ", "") + units = { + "B": 1, + "KB": 1000, + "MB": 1000**2, + "GB": 1000**3, + "TB": 1000**4, + "KIB": 1024, + "MIB": 1024**2, + "GIB": 1024**3, + "TIB": 1024**4, + } + for suffix in sorted(units, key=len, reverse=True): + if text.endswith(suffix): + number = text[: -len(suffix)] + if not number: + raise argparse.ArgumentTypeError(f"invalid size: {value}") + size = int(float(number) * units[suffix]) + if size <= 0: + raise argparse.ArgumentTypeError(f"size must be positive: {value}") + return size + try: + size = int(text) + except ValueError as exc: + raise argparse.ArgumentTypeError(f"invalid size: {value}") from exc + if size <= 0: + raise argparse.ArgumentTypeError(f"size must be positive: {value}") + return size + + +def _to_device(tensor: torch.Tensor, device: torch.device) -> torch.Tensor: + return tensor.to(device=device, non_blocking=device.type == "cuda") + + +def dequantize_fp8_block( + weight_slice, + scale: torch.Tensor, + *, + output_shape: Sequence[int], + block_size: int, + row_chunk_size: int, + device: torch.device, +) -> torch.Tensor: + """Dequantize a 2-D E4M3 matrix with a 2-D E8M0 block scale.""" + + rows, cols = (int(dim) for dim in output_shape) + expected_scale_shape = (math.ceil(rows / block_size), math.ceil(cols / block_size)) + if tuple(scale.shape) != expected_scale_shape: + raise ValueError( + f"FP8 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}" + ) + + output = torch.empty((rows, cols), dtype=torch.bfloat16, device="cpu") + for row_start in range(0, rows, row_chunk_size): + row_end = min(row_start + row_chunk_size, rows) + scale_row_start = row_start // block_size + scale_row_end = math.ceil(row_end / block_size) + scale_row_offset = row_start - scale_row_start * block_size + + weight_chunk = _to_device(weight_slice[row_start:row_end], device).to( + torch.float32 + ) + scale_chunk = _to_device(scale[scale_row_start:scale_row_end], device).to( + torch.float32 + ) + expanded_scale = scale_chunk.repeat_interleave(block_size, dim=0) + expanded_scale = expanded_scale[ + scale_row_offset : scale_row_offset + row_end - row_start + ].repeat_interleave(block_size, dim=1)[:, :cols] + converted = (weight_chunk * expanded_scale).to(torch.bfloat16).cpu() + output[row_start:row_end].copy_(converted) + del weight_chunk, scale_chunk, expanded_scale, converted + return output + + +def dequantize_mxfp4( + packed_weight_slice, + scale: torch.Tensor, + *, + packed_shape: Sequence[int], + block_size: int, + row_chunk_size: int, + device: torch.device, +) -> torch.Tensor: + """Unpack a 2-D MXFP4 matrix and apply its per-K-block E8M0 scale.""" + + rows, packed_cols = (int(dim) for dim in packed_shape) + logical_cols = packed_cols * 2 + if logical_cols % block_size != 0: + raise ValueError( + f"MXFP4 logical K dimension {logical_cols} is not divisible by block size {block_size}" + ) + expected_scale_shape = (rows, logical_cols // block_size) + if tuple(scale.shape) != expected_scale_shape: + raise ValueError( + f"MXFP4 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}" + ) + + output = torch.empty((rows, logical_cols), dtype=torch.bfloat16, device="cpu") + lookup = torch.tensor(FP4_VALUES, dtype=torch.float32, device=device) + for row_start in range(0, rows, row_chunk_size): + row_end = min(row_start + row_chunk_size, rows) + packed = _to_device(packed_weight_slice[row_start:row_end], device).view( + torch.uint8 + ) + low = (packed & 0x0F).to(torch.long) + high = (packed >> 4).to(torch.long) + + values = torch.empty( + (row_end - row_start, logical_cols), dtype=torch.float32, device=device + ) + values[:, 0::2] = lookup[low] + values[:, 1::2] = lookup[high] + scale_chunk = _to_device(scale[row_start:row_end], device).to(torch.float32) + expanded_scale = scale_chunk.repeat_interleave(block_size, dim=1) + converted = (values * expanded_scale).to(torch.bfloat16).cpu() + output[row_start:row_end].copy_(converted) + del packed, low, high, values, scale_chunk, expanded_scale, converted + return output + + +def _load_index(model_dir: Path) -> Tuple[Path, dict]: + index_path = model_dir / "model.safetensors.index.json" + if not index_path.is_file(): + raise FileNotFoundError(f"missing safetensors index: {index_path}") + with index_path.open("r", encoding="utf-8") as file: + index = json.load(file) + if not isinstance(index.get("weight_map"), dict): + raise ValueError(f"invalid weight_map in {index_path}") + return index_path, index + + +def _inspect_specs( + model_dir: Path, + weight_map: Mapping[str, str], + *, + prefix: str, + fp8_block_size: int, + mxfp4_block_size: int, +) -> Tuple[List[TensorSpec], set[str], int]: + mtp_weight_map = { + name: shard for name, shard in weight_map.items() if name.startswith(prefix) + } + if not mtp_weight_map: + raise ValueError(f"no checkpoint tensors start with {prefix!r}") + + names_by_shard: Dict[str, List[str]] = {} + for name, shard in mtp_weight_map.items(): + names_by_shard.setdefault(shard, []).append(name) + + tensor_meta: Dict[str, Tuple[str, Tuple[int, ...]]] = {} + original_mtp_nbytes = 0 + for shard, names in names_by_shard.items(): + shard_path = model_dir / shard + if not shard_path.is_file(): + raise FileNotFoundError(f"missing shard: {shard_path}") + with safe_open(shard_path, framework="pt", device="cpu") as file: + shard_keys = set(file.keys()) + missing = set(names) - shard_keys + if missing: + raise ValueError( + f"{shard} is missing indexed tensors: {sorted(missing)[:5]}" + ) + for name in names: + tensor_slice = file.get_slice(name) + dtype = tensor_slice.get_dtype() + shape = tuple(int(dim) for dim in tensor_slice.get_shape()) + tensor_meta[name] = (dtype, shape) + original_mtp_nbytes += _tensor_nbytes(dtype, shape) + + paired_scales: set[str] = set() + specs: List[TensorSpec] = [] + for name in mtp_weight_map: + dtype, shape = tensor_meta[name] + if dtype == "F8_E8M0": + continue + if dtype == "F8_E4M3": + if len(shape) != 2: + raise ValueError(f"FP8 tensor must be 2-D: {name} has shape {shape}") + scale_name = _scale_name(name) + if tensor_meta.get(scale_name, (None,))[0] != "F8_E8M0": + raise ValueError( + f"missing E8M0 scale for {name}: expected {scale_name}" + ) + scale_shape = tensor_meta[scale_name][1] + expected_scale_shape = ( + math.ceil(shape[0] / fp8_block_size), + math.ceil(shape[1] / fp8_block_size), + ) + if scale_shape != expected_scale_shape: + raise ValueError( + f"FP8 scale shape mismatch for {name}: expected {expected_scale_shape}, got {scale_shape}" + ) + paired_scales.add(scale_name) + output_shape = shape + conversion = "fp8_block" + elif dtype == "I8": + if len(shape) != 2: + raise ValueError( + f"packed MXFP4 tensor must be 2-D: {name} has shape {shape}" + ) + scale_name = _scale_name(name) + if tensor_meta.get(scale_name, (None,))[0] != "F8_E8M0": + raise ValueError( + f"missing E8M0 scale for {name}: expected {scale_name}" + ) + output_shape = (shape[0], shape[1] * 2) + expected_scale_shape = ( + output_shape[0], + output_shape[1] // mxfp4_block_size, + ) + if tensor_meta[scale_name][1] != expected_scale_shape: + raise ValueError( + f"MXFP4 scale shape mismatch for {name}: expected {expected_scale_shape}, " + f"got {tensor_meta[scale_name][1]}" + ) + paired_scales.add(scale_name) + conversion = "mxfp4" + elif dtype in {"BF16", "F16", "F32"}: + scale_name = None + output_shape = shape + conversion = "cast" + else: + raise ValueError( + f"cannot convert {name} with dtype {dtype} to BF16 without model-specific semantics" + ) + + specs.append( + TensorSpec( + name=name, + source_shard=mtp_weight_map[name], + source_dtype=dtype, + source_shape=shape, + output_shape=output_shape, + output_nbytes=_numel(output_shape) * 2, + conversion=conversion, + scale_name=scale_name, + ) + ) + + orphan_scales = { + name for name, (dtype, _) in tensor_meta.items() if dtype == "F8_E8M0" + } - paired_scales + if orphan_scales: + raise ValueError(f"orphan MTP E8M0 scales: {sorted(orphan_scales)[:10]}") + return specs, paired_scales, original_mtp_nbytes + + +def _plan_shards( + specs: Sequence[TensorSpec], max_shard_size: int +) -> List[List[TensorSpec]]: + shards: List[List[TensorSpec]] = [] + current: List[TensorSpec] = [] + current_size = 0 + for spec in specs: + if current and current_size + spec.output_nbytes > max_shard_size: + shards.append(current) + current = [] + current_size = 0 + current.append(spec) + current_size += spec.output_nbytes + if current: + shards.append(current) + return shards + + +def _copy_or_link(source: Path, destination: Path, mode: str) -> None: + if mode == "copy": + shutil.copy2(source, destination) + return + if mode == "symlink": + destination.symlink_to(source.resolve()) + return + try: + os.link(source, destination) + except OSError: + shutil.copy2(source, destination) + + +def _prepare_output_dir( + source_dir: Path, + output_dir: Path, + *, + referenced_non_mtp_shards: Iterable[str], + link_mode: str, +) -> None: + if output_dir.exists(): + if any(output_dir.iterdir()): + raise FileExistsError(f"output directory is not empty: {output_dir}") + else: + output_dir.mkdir(parents=True) + + for entry in source_dir.iterdir(): + if not entry.is_file() or entry.name == "model.safetensors.index.json": + continue + if entry.suffix == ".safetensors": + continue + _copy_or_link(entry, output_dir / entry.name, link_mode) + + for shard in sorted(set(referenced_non_mtp_shards)): + _copy_or_link(source_dir / shard, output_dir / shard, link_mode) + + +def _convert_spec( + spec: TensorSpec, + source_file, + *, + fp8_block_size: int, + mxfp4_block_size: int, + row_chunk_size: int, + device: torch.device, +) -> torch.Tensor: + if spec.conversion == "cast": + return source_file.get_tensor(spec.name).to(torch.bfloat16).contiguous() + + assert spec.scale_name is not None + weight_slice = source_file.get_slice(spec.name) + scale = source_file.get_tensor(spec.scale_name) + if spec.conversion == "fp8_block": + return dequantize_fp8_block( + weight_slice, + scale, + output_shape=spec.output_shape, + block_size=fp8_block_size, + row_chunk_size=row_chunk_size, + device=device, + ) + if spec.conversion == "mxfp4": + return dequantize_mxfp4( + weight_slice, + scale, + packed_shape=spec.source_shape, + block_size=mxfp4_block_size, + row_chunk_size=row_chunk_size, + device=device, + ) + raise AssertionError(f"unknown conversion: {spec.conversion}") + + +def _write_converted_shards( + source_dir: Path, + output_dir: Path, + shard_plan: Sequence[Sequence[TensorSpec]], + *, + fp8_block_size: int, + mxfp4_block_size: int, + row_chunk_size: int, + device: torch.device, +) -> Dict[str, str]: + source_shards = sorted( + {spec.source_shard for shard in shard_plan for spec in shard} + ) + output_weight_map: Dict[str, str] = {} + with ExitStack() as stack: + source_files = { + shard: stack.enter_context( + safe_open(source_dir / shard, framework="pt", device="cpu") + ) + for shard in source_shards + } + shard_count = len(shard_plan) + for shard_index, shard_specs in enumerate(shard_plan, start=1): + output_name = f"mtp-bf16-{shard_index:05d}-of-{shard_count:05d}.safetensors" + planned_bytes = sum(spec.output_nbytes for spec in shard_specs) + print( + f"[{shard_index}/{shard_count}] converting {len(shard_specs)} tensors " + f"({planned_bytes / 1024**3:.2f} GiB) -> {output_name}", + flush=True, + ) + tensors = { + spec.name: _convert_spec( + spec, + source_files[spec.source_shard], + fp8_block_size=fp8_block_size, + mxfp4_block_size=mxfp4_block_size, + row_chunk_size=row_chunk_size, + device=device, + ) + for spec in shard_specs + } + temporary_path = output_dir / f".{output_name}.tmp" + output_path = output_dir / output_name + save_file(tensors, temporary_path, metadata={"format": "pt"}) + os.replace(temporary_path, output_path) + for spec in shard_specs: + output_weight_map[spec.name] = output_name + del tensors + gc.collect() + if device.type == "cuda": + torch.cuda.empty_cache() + return output_weight_map + + +def _verify_output( + output_dir: Path, + specs: Sequence[TensorSpec], + converted_weight_map: Mapping[str, str], +) -> None: + expected_by_shard: Dict[str, Dict[str, TensorSpec]] = {} + for spec in specs: + expected_by_shard.setdefault(converted_weight_map[spec.name], {})[ + spec.name + ] = spec + + for shard, expected in expected_by_shard.items(): + with safe_open(output_dir / shard, framework="pt", device="cpu") as file: + actual_keys = set(file.keys()) + if actual_keys != set(expected): + raise ValueError( + f"output key mismatch in {shard}: missing={set(expected) - actual_keys}, " + f"unexpected={actual_keys - set(expected)}" + ) + for name, spec in expected.items(): + tensor_slice = file.get_slice(name) + if tensor_slice.get_dtype() != "BF16": + raise ValueError( + f"{name} is {tensor_slice.get_dtype()}, expected BF16" + ) + if tuple(tensor_slice.get_shape()) != spec.output_shape: + raise ValueError( + f"{name} shape is {tuple(tensor_slice.get_shape())}, expected {spec.output_shape}" + ) + + +def convert(args: argparse.Namespace) -> None: + source_dir = Path(args.source_model_dir).expanduser().resolve(strict=True) + output_dir = Path(args.output_model_dir).expanduser().resolve(strict=False) + if source_dir == output_dir: + raise ValueError("source and output directories must be different") + if args.row_chunk_size <= 0: + raise ValueError("--row-chunk-size must be positive") + + device = torch.device(args.device) + if device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError( + "CUDA conversion requested but torch.cuda.is_available() is false" + ) + + _, index = _load_index(source_dir) + weight_map: Dict[str, str] = index["weight_map"] + specs, paired_scales, original_mtp_nbytes = _inspect_specs( + source_dir, + weight_map, + prefix=args.prefix, + fp8_block_size=args.fp8_block_size, + mxfp4_block_size=args.mxfp4_block_size, + ) + shard_plan = _plan_shards(specs, args.max_shard_size) + output_mtp_nbytes = sum(spec.output_nbytes for spec in specs) + conversion_counts: Dict[str, int] = {} + for spec in specs: + conversion_counts[spec.conversion] = ( + conversion_counts.get(spec.conversion, 0) + 1 + ) + + print(f"source: {source_dir}") + print(f"output: {output_dir}") + print(f"MTP output tensors: {len(specs)}; removed scales: {len(paired_scales)}") + print(f"conversions: {conversion_counts}") + print( + f"MTP size: {original_mtp_nbytes / 1024**3:.2f} GiB -> " + f"{output_mtp_nbytes / 1024**3:.2f} GiB in {len(shard_plan)} shards" + ) + if args.dry_run: + return + + non_mtp_weight_map = { + name: shard + for name, shard in weight_map.items() + if not name.startswith(args.prefix) + } + _prepare_output_dir( + source_dir, + output_dir, + referenced_non_mtp_shards=non_mtp_weight_map.values(), + link_mode=args.link_mode, + ) + converted_weight_map = _write_converted_shards( + source_dir, + output_dir, + shard_plan, + fp8_block_size=args.fp8_block_size, + mxfp4_block_size=args.mxfp4_block_size, + row_chunk_size=args.row_chunk_size, + device=device, + ) + + new_weight_map: Dict[str, str] = {} + for name, shard in weight_map.items(): + if not name.startswith(args.prefix): + new_weight_map[name] = shard + elif name in converted_weight_map: + new_weight_map[name] = converted_weight_map[name] + elif name not in paired_scales: + raise AssertionError( + f"MTP tensor was neither converted nor removed: {name}" + ) + + metadata = dict(index.get("metadata") or {}) + if "total_size" in metadata: + metadata["total_size"] = ( + int(metadata["total_size"]) - original_mtp_nbytes + output_mtp_nbytes + ) + new_index = {"metadata": metadata, "weight_map": new_weight_map} + index_path = output_dir / "model.safetensors.index.json" + temporary_index_path = output_dir / ".model.safetensors.index.json.tmp" + with temporary_index_path.open("w", encoding="utf-8") as file: + json.dump(new_index, file, ensure_ascii=False, indent=2) + file.write("\n") + os.replace(temporary_index_path, index_path) + + manifest = { + "source_model_dir": str(source_dir), + "weight_prefix": args.prefix, + "output_dtype": "bfloat16", + "runtime_note": ( + "Configure converted MTP layers with quant type 'none'; target layers " + "remain on their original quantization path." + ), + "fp8_block_size": args.fp8_block_size, + "mxfp4_block_size": args.mxfp4_block_size, + "converted_tensor_count": len(specs), + "removed_scale_count": len(paired_scales), + "original_mtp_bytes": original_mtp_nbytes, + "output_mtp_bytes": output_mtp_nbytes, + } + with (output_dir / "mtp_bf16_conversion.json").open("w", encoding="utf-8") as file: + json.dump(manifest, file, ensure_ascii=False, indent=2) + file.write("\n") + + if not args.no_verify: + print("verifying converted shards ...", flush=True) + _verify_output(output_dir, specs, converted_weight_map) + print(f"conversion complete: {output_dir}") + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("source_model_dir", help="source Hugging Face model directory") + parser.add_argument("output_model_dir", help="new output model directory") + parser.add_argument( + "--prefix", default="mtp.", help="checkpoint key prefix to convert" + ) + parser.add_argument( + "--device", + default="cpu", + help="conversion device, for example cpu or cuda:0 (output shards are always stored on CPU)", + ) + parser.add_argument( + "--max-shard-size", + type=_parse_size, + default=_parse_size("2GiB"), + help="maximum planned size of each converted shard (default: 2GiB)", + ) + parser.add_argument( + "--row-chunk-size", + type=int, + default=256, + help="number of matrix rows converted at once (default: 256)", + ) + parser.add_argument("--fp8-block-size", type=int, default=128) + parser.add_argument("--mxfp4-block-size", type=int, default=32) + parser.add_argument( + "--link-mode", + choices=("hardlink", "copy", "symlink"), + default="hardlink", + help="how unchanged model files are placed in the output directory", + ) + parser.add_argument( + "--dry-run", + action="store_true", + help="inspect and print the conversion plan only", + ) + parser.add_argument( + "--no-verify", + action="store_true", + help="skip the final dtype/key/shape verification", + ) + return parser + + +def main() -> None: + convert(build_parser().parse_args()) + + +if __name__ == "__main__": + main() From 5dfa35dbb5005721b48e99ff18ebc1d5f1d1b1ae Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:37:09 +0800 Subject: [PATCH 139/214] fix: pad DSpark scratch pages for CUDA graphs --- lightllm/models/deepseek_v4_dspark/model.py | 26 +++++++++++++ test/unit/test_deepseek_v4_dspark.py | 43 +++++++++++++++++++++ 2 files changed, 69 insertions(+) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index aaf4765e1a..1dea813f3b 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -3,6 +3,7 @@ from typing import List import torch +import torch.nn.functional as F from safetensors import safe_open from tqdm import tqdm @@ -127,6 +128,31 @@ def _init_req_manager(self): def _init_mem_manager(self): self.mem_manager = self.main_model.mem_manager + def _create_padded_decode_model_input( + self, model_input: ModelInput, new_batch_size: int + ): + padded_input = super()._create_padded_decode_model_input( + model_input, new_batch_size + ) + scratch_pages = model_input.mtp_draft_swa_pages + if padded_input is model_input or scratch_pages is None: + return padded_input + + assert model_input.batch_size % self.block_size == 0 + assert new_batch_size % self.block_size == 0 + page_num = model_input.batch_size // self.block_size + padded_page_num = new_batch_size // self.block_size + assert scratch_pages.shape == (page_num,) + # HOLD rows ignore the page value. Pad only the CUDA Graph input; the CPU + # owner list must continue to contain exactly the pages that were allocated. + padded_input.mtp_draft_swa_pages = F.pad( + scratch_pages, + (0, padded_page_num - page_num), + mode="constant", + value=0, + ) + return padded_input + def _init_infer_layer(self, start_layer_index=None): assert start_layer_index is None self.pre_infer = self.pre_layer_infer_class(network_config=self.config) diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index ec06ac9dec..cbf65705c7 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -34,6 +34,49 @@ def test_deepseek_v4_dspark_extends_target_compress_rates_for_draft_layers(): assert rates == target_rates + [0, 0, 0] +def test_dspark_cuda_graph_padding_extends_only_gpu_scratch_pages(): + from lightllm.common.basemodel.batch_objs import ModelInput + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + block_size = 5 + batch_size = 14 * block_size + graph_batch_size = 16 * block_size + model = DeepseekV4DSparkModel.__new__(DeepseekV4DSparkModel) + model.block_size = block_size + model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=127) + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=255) + pages_cpu = torch.arange(14, dtype=torch.int32) + model_input = ModelInput( + batch_size=batch_size, + total_token_num=batch_size, + max_q_seq_len=1, + max_kv_seq_len=8, + input_ids=torch.ones(batch_size, dtype=torch.int64), + b_req_idx=torch.arange(batch_size, dtype=torch.int32), + b_mtp_index=torch.zeros(batch_size, dtype=torch.int32), + b_seq_len=torch.full((batch_size,), 8, dtype=torch.int32), + b_position_delta=torch.zeros(batch_size, dtype=torch.int32), + b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32), + b_shared_radix_node_id=torch.full((batch_size,), -1, dtype=torch.int64), + mem_indexes=torch.arange(batch_size, dtype=torch.int32), + is_prefill=False, + multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], + mtp_draft_swa_pages_cpu=pages_cpu, + mtp_draft_swa_pages=pages_cpu.clone(), + ) + + padded_input = model._create_padded_decode_model_input( + model_input, graph_batch_size + ) + + assert padded_input.mtp_draft_swa_pages.shape == (16,) + torch.testing.assert_close(padded_input.mtp_draft_swa_pages[:14], pages_cpu) + torch.testing.assert_close( + padded_input.mtp_draft_swa_pages[14:], torch.zeros(2, dtype=torch.int32) + ) + assert padded_input.mtp_draft_swa_pages_cpu is pages_cpu + + def test_dspark_scratch_cleanup_keeps_page_ids_on_cpu(): from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( From f0e36f3fe50bd7ab6aaee4762176f15bda1457e8 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 15:55:05 +0800 Subject: [PATCH 140/214] feat: support DSpark in DP backend --- .../mode_backend/dp_backend/impl.py | 18 ++++++++++++------ .../mtp_speculative/proposers/dspark.py | 7 ++++--- .../test_dp_overlap_spec_engine.py | 17 +++++++++++++++++ 3 files changed, 33 insertions(+), 9 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py index ade78f9aae..d9a28bdbfc 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 @@ -33,8 +33,12 @@ def __init__(self) -> None: # 在 mtp 模式下切换绑定的prefill 和 decode 函数 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 spec_mode == "dflash": + raise NotImplementedError("DP backend does not support DFlash parallel block drafting yet.") + if spec_mode == "dspark" and ( + self.enable_prefill_microbatch_overlap or self.enable_decode_microbatch_overlap + ): + raise NotImplementedError("DP DSpark does not support prefill/decode microbatch overlap yet.") if self.enable_prefill_microbatch_overlap: self.prefill = self.prefill_overlap_mtp else: @@ -70,10 +74,12 @@ def init_spec_engine(self): enable_dynmaic_mtp=self.args.mtp_dynamic_verify, ) - self.dp_overlap_spec_engine = DPOverlapSpecEngine( - **engine_kwargs, - common_engine=self.spec_engine, - ) + self.dp_overlap_spec_engine = None + if self.enable_prefill_microbatch_overlap or self.enable_decode_microbatch_overlap: + 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 ) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 3144006f60..26d0767f7f 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -90,9 +90,10 @@ def propose_next( # 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) + if req_num > 0: + 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,其余行 diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index fcdf4e3295..54690186e8 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -57,6 +57,23 @@ def test_dp_backend_reuses_common_engine_outside_overlap(): assert not issubclass(DPOverlapSpecEngine, SpecEngine) +def test_dp_backend_dspark_uses_common_engine_without_overlap(): + backend = DPChunkedPrefillBackend.__new__(DPChunkedPrefillBackend) + backend.args = SimpleNamespace(mtp_mode="dspark", mtp_dynamic_verify=False, dp=2) + backend.max_draft_step = 5 + backend.dp_size = 2 + backend.draft_models = [SimpleNamespace(block_size=5, graph=None)] + backend.enable_prefill_microbatch_overlap = False + backend.enable_decode_microbatch_overlap = False + + backend.init_spec_engine() + + assert type(backend.spec_engine) is SpecEngine + assert backend.dp_overlap_spec_engine is None + assert backend.prefill_draft_engine is backend.spec_engine + assert backend.decode_draft_engine is backend.spec_engine + + def test_lightspec_planner_reduces_draft_step_with_max(monkeypatch): group = object() planner = LightSpecPlanner.__new__(LightSpecPlanner) From 4afca8c89510c0e8cac70c249092ddc8bbd17231 Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 16:58:16 +0800 Subject: [PATCH 141/214] fix: keep dspark empty DP ranks in sync --- lightllm/models/deepseek_v4_dspark/model.py | 11 +++++- .../mtp_speculative/proposers/dspark.py | 6 +-- test/unit/test_deepseek_v4_dspark.py | 37 +++++++++++++++++++ .../mtp_speculative/test_dspark.py | 30 +++++++++++++++ 4 files changed, 80 insertions(+), 4 deletions(-) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 1dea813f3b..9dbf2888ae 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -131,6 +131,16 @@ def _init_mem_manager(self): def _create_padded_decode_model_input( self, model_input: ModelInput, new_batch_size: int ): + # A DSpark logical request always occupies one complete physical block. + # Keep the generic HOLD padding semantics, but make its requested size + # valid even for an empty eager-mode DP rank (0 -> one HOLD block). + assert model_input.batch_size <= new_batch_size + assert model_input.batch_size % self.block_size == 0 + new_batch_size = max(new_batch_size, self.block_size) + new_batch_size = ( + (new_batch_size + self.block_size - 1) // self.block_size + ) * self.block_size + padded_input = super()._create_padded_decode_model_input( model_input, new_batch_size ) @@ -138,7 +148,6 @@ def _create_padded_decode_model_input( if padded_input is model_input or scratch_pages is None: return padded_input - assert model_input.batch_size % self.block_size == 0 assert new_batch_size % self.block_size == 0 page_num = model_input.batch_size // self.block_size padded_page_num = new_batch_size // self.block_size diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 26d0767f7f..a572d80bc9 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -38,11 +38,11 @@ def fill_draft_model_kv_state( 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 + # KV 写入唯一的 parallel-block draft model。空 DP rank 也必须进入 + # draft forward,由通用 HOLD padding 补出 dummy token,保证各 rank + # 执行相同次数的 DeepEP collective。使用浅副本,避免把 draft # 专用 hidden 挂到后续流程仍可能读取的 target ModelInput 上。 draft_input = copy.copy(target_model_input) draft_input.mtp_draft_input_hiddens = target_hidden diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index cbf65705c7..861b4819e2 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -77,6 +77,43 @@ def test_dspark_cuda_graph_padding_extends_only_gpu_scratch_pages(): assert padded_input.mtp_draft_swa_pages_cpu is pages_cpu +def test_dspark_empty_decode_padding_builds_one_hold_block(): + from lightllm.common.basemodel.batch_objs import ModelInput + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + block_size = 5 + model = DeepseekV4DSparkModel.__new__(DeepseekV4DSparkModel) + model.block_size = block_size + model.req_manager = SimpleNamespace(HOLD_REQUEST_ID=127) + model.mem_manager = SimpleNamespace(HOLD_TOKEN_MEMINDEX=255) + 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=[], + ) + + padded_input = model._create_padded_decode_model_input(model_input, 1) + + assert model_input.batch_size == 0 + assert padded_input.batch_size == block_size + assert padded_input.total_token_num == 2 * block_size + assert padded_input.input_ids.tolist() == [1] * block_size + assert padded_input.b_req_idx.tolist() == [127] * block_size + assert padded_input.b_seq_len.tolist() == [2] * block_size + assert padded_input.mem_indexes.tolist() == [255] * block_size + + def test_dspark_scratch_cleanup_keeps_page_ids_on_cpu(): from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.server.router.model_infer.mtp_speculative.proposers.base import ( 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 index d341dc7dbd..7f42f8ecc6 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py @@ -37,6 +37,36 @@ def test_dspark_prefill_uses_a_shallow_copy_for_target_hidden(): assert model_input.mtp_draft_input_hiddens is None +def test_dspark_empty_prefill_still_forwards_for_dp_collectives(): + 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.empty((0,), dtype=torch.int32), + input_ids=torch.empty((0,), dtype=torch.int64), + mtp_draft_input_hiddens=None, + ) + target_hidden = torch.empty((0, 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.empty((0,), 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) From 53bba62886b9ac27d45361f53ac0f0480c0c34d5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 27 Aug 2026 10:06:08 +0000 Subject: [PATCH 142/214] delte == --- lightllm/models/deepseek_v4_dspark/model.py | 64 +++++---------------- 1 file changed, 15 insertions(+), 49 deletions(-) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 9dbf2888ae..8e137cd351 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -55,20 +55,14 @@ def __init__(self, kvargs: dict): 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" - ) - assert ( - not self.mtp_previous_draft_models - ), "DeepSeek-V4 DSpark uses one parallel draft model" + self.mtp_previous_draft_models: List[TpPartBaseModel] = kvargs.pop("mtp_previous_draft_models") + assert not self.mtp_previous_draft_models, "DeepSeek-V4 DSpark uses one parallel draft model" def _init_config(self): super()._init_config() draft_layer_num = len(self.config["compress_ratios"]) - self.config["n_layer"] assert draft_layer_num > 0 - assert ( - self.config["compress_ratios"][-draft_layer_num:] == [0] * draft_layer_num - ) + assert self.config["compress_ratios"][-draft_layer_num:] == [0] * draft_layer_num assert len(self.config["dspark_target_layer_ids"]) == draft_layer_num self.config["dspark_layer_num"] = draft_layer_num @@ -82,22 +76,12 @@ def _init_config(self): def _verify_params(self): super()._verify_params() - assert ( - not self.enable_tpsp_mix_mode - ), "DeepSeek-V4 DSpark draft model does not support TP-SP" - assert self.args.mtp_step == self.config["dspark_block_size"], ( - f"DeepSeek-V4 DSpark requires --mtp_step {self.config['dspark_block_size']}, " - f"got {self.args.mtp_step}" - ) + assert not self.enable_tpsp_mix_mode, "DeepSeek-V4 DSpark draft model does not support TP-SP" def _init_quant(self): super()._init_quant() - expert_quant_type = self.quant_cfg.get_quant_type( - self.config["n_layer"] - 1, "fused_moe" - ) - for layer_idx in range( - self.config["n_layer"], len(self.config["compress_ratios"]) - ): + expert_quant_type = self.quant_cfg.get_quant_type(self.config["n_layer"] - 1, "fused_moe") + for layer_idx in range(self.config["n_layer"], len(self.config["compress_ratios"])): self.quant_cfg.quant_cfg[layer_idx]["fused_moe"] = expert_quant_type def _init_weights(self, start_layer_index=None): @@ -108,9 +92,7 @@ def _init_weights(self, start_layer_index=None): quant_cfg=self.quant_cfg, ) self.pre_post_weight.wte_weight_ = self.main_model.pre_post_weight.wte_weight_ - self.pre_post_weight.lm_head_weight_ = ( - self.main_model.pre_post_weight.lm_head_weight_ - ) + self.pre_post_weight.lm_head_weight_ = self.main_model.pre_post_weight.lm_head_weight_ first_layer = self.config["n_layer"] self.trans_layers_weight = [ self.transformer_weight_class( @@ -128,22 +110,16 @@ def _init_req_manager(self): def _init_mem_manager(self): self.mem_manager = self.main_model.mem_manager - def _create_padded_decode_model_input( - self, model_input: ModelInput, new_batch_size: int - ): + def _create_padded_decode_model_input(self, model_input: ModelInput, new_batch_size: int): # A DSpark logical request always occupies one complete physical block. # Keep the generic HOLD padding semantics, but make its requested size # valid even for an empty eager-mode DP rank (0 -> one HOLD block). assert model_input.batch_size <= new_batch_size assert model_input.batch_size % self.block_size == 0 new_batch_size = max(new_batch_size, self.block_size) - new_batch_size = ( - (new_batch_size + self.block_size - 1) // self.block_size - ) * self.block_size + new_batch_size = ((new_batch_size + self.block_size - 1) // self.block_size) * self.block_size - padded_input = super()._create_padded_decode_model_input( - model_input, new_batch_size - ) + padded_input = super()._create_padded_decode_model_input(model_input, new_batch_size) scratch_pages = model_input.mtp_draft_swa_pages if padded_input is model_input or scratch_pages is None: return padded_input @@ -168,9 +144,7 @@ def _init_infer_layer(self, start_layer_index=None): self.post_infer = self.post_layer_infer_class(network_config=self.config) first_layer = self.config["n_layer"] self.layers_infer = [ - self.transformer_layer_infer_class( - first_layer + stage_id, network_config=self.config - ) + self.transformer_layer_infer_class(first_layer + stage_id, network_config=self.config) for stage_id in range(self.config["dspark_layer_num"]) ] @@ -233,9 +207,7 @@ def _gen_special_model_input(self, token_num: int): # CUDA Graph capture must take the same direct scratch-slot branch # as runtime. HOLD rows never dereference the page value, but the # tensor gives the captured graph stable storage for replay copies. - "mtp_draft_swa_pages": torch.zeros( - token_num // self.block_size, dtype=torch.int32, device="cuda" - ), + "mtp_draft_swa_pages": torch.zeros(token_num // self.block_size, dtype=torch.int32, device="cuda"), } def _autotune_warmup(self): @@ -252,16 +224,10 @@ def load_weights(self, weight_dict: dict): return super().load_weights(weight_dict) index_file = os.path.join(self.weight_dir_, "model.safetensors.index.json") - assert utils.PetrelHelper.exists( - index_file - ), "DeepSeek-V4 DSpark requires model.safetensors.index.json" + assert utils.PetrelHelper.exists(index_file), "DeepSeek-V4 DSpark requires model.safetensors.index.json" weight_map = utils.PetrelHelper.load_json(index_file)["weight_map"] - candidate_files = sorted( - {file_ for key, file_ in weight_map.items() if key.startswith("mtp.")} - ) - assert ( - candidate_files - ), "DeepSeek-V4 DSpark weights with prefix mtp.* were not found" + candidate_files = sorted({file_ for key, file_ in weight_map.items() if key.startswith("mtp.")}) + assert candidate_files, "DeepSeek-V4 DSpark weights with prefix mtp.* were not found" loaded_key_count = 0 desc = f"pid {os.getpid()} Loading DeepSeek-V4 DSpark weights" From 8c9942a56621300a01e454a28a6f60d71865f09a Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Thu, 27 Aug 2026 18:32:28 +0800 Subject: [PATCH 143/214] fix: use mtp step for DeepSeek-V4 DSpark width --- lightllm/models/deepseek_v4_dspark/model.py | 10 ++++++- test/unit/test_deepseek_v4_dspark.py | 31 +++++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 8e137cd351..f483b54645 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -77,6 +77,14 @@ def _init_config(self): def _verify_params(self): super()._verify_params() assert not self.enable_tpsp_mix_mode, "DeepSeek-V4 DSpark draft model does not support TP-SP" + checkpoint_block_size = int(self.config["dspark_block_size"]) + assert 0 < self.args.mtp_step <= checkpoint_block_size, ( + f"DeepSeek-V4 DSpark requires --mtp_step in [1, {checkpoint_block_size}], " + f"got {self.args.mtp_step}" + ) + # The checkpoint block size is its maximum supported width. Runtime + # allocation, indexing, and Markov decoding follow --mtp_step. + self.config["block_size"] = self.args.mtp_step def _init_quant(self): super()._init_quant() @@ -160,7 +168,7 @@ def _init_custom(self): self._cos_cached_compress = self.main_model._cos_cached_compress self._sin_cached_compress = self.main_model._sin_cached_compress self.dsv4_workspace = self.main_model.dsv4_workspace - self.block_size = self.config["dspark_block_size"] + self.block_size = int(self.config["block_size"]) self.mask_token_id = self.config["dspark_noise_token_id"] for layer in self.layers_infer: layer.freqs_cis = self._freqs_cis_sliding diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index 861b4819e2..f972e46a6e 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -8,6 +8,37 @@ from lightllm.utils.config_utils import get_deepseek_v4_compress_rates +@pytest.mark.parametrize("mtp_step", [1, 4, 5]) +def test_deepseek_v4_dspark_runtime_width_follows_mtp_step(monkeypatch, mtp_step): + from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + monkeypatch.setattr(DeepseekV4TpPartModel, "_verify_params", lambda self: None) + model = DeepseekV4DSparkModel.__new__(DeepseekV4DSparkModel) + model.args = SimpleNamespace(mtp_step=mtp_step) + model.config = {"block_size": 5, "dspark_block_size": 5} + model.enable_tpsp_mix_mode = False + + model._verify_params() + + assert model.config["block_size"] == mtp_step + + +@pytest.mark.parametrize("mtp_step", [0, 6]) +def test_deepseek_v4_dspark_rejects_invalid_runtime_width(monkeypatch, mtp_step): + from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + monkeypatch.setattr(DeepseekV4TpPartModel, "_verify_params", lambda self: None) + model = DeepseekV4DSparkModel.__new__(DeepseekV4DSparkModel) + model.args = SimpleNamespace(mtp_step=mtp_step) + model.config = {"block_size": 5, "dspark_block_size": 5} + model.enable_tpsp_mix_mode = False + + with pytest.raises(AssertionError, match=r"requires --mtp_step in \[1, 5\]"): + model._verify_params() + + def test_deepseek_v4_dspark_layer_count_uses_trailing_swa_layers(tmp_path): config = { "model_type": "deepseek_v4", From 811230e0755f2922cf18c63dcfedfe75453da758 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 27 Aug 2026 12:33:43 +0000 Subject: [PATCH 144/214] fix ZMQ crash when cache workers finish concurrently --- .../server/multi_level_kv_cache/manager.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index ef5b7369c9..4f8a6d007c 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -9,7 +9,7 @@ import threading import concurrent.futures import setproctitle -from queue import Queue +from queue import Empty, Queue from typing import List from lightllm.server.core.objs import ShmReqManager, Req, StartArgs from lightllm.server.core.objs.io_objs import GroupReqIndexes @@ -45,6 +45,9 @@ def __init__( # 控制进行 cpu cache 页面匹配的时间,超过时间则不再匹配,直接转发。 self.cpu_cache_time_out = 0.5 self.recv_queue = Queue(maxsize=1024) + # ZeroMQ sockets are not thread-safe. Cache workers enqueue completed + # requests so the recv_loop thread remains the sole socket owner. + self.send_to_router_queue = Queue() self.cpu_cache_thread = threading.Thread(target=self.cpu_cache_hanle_loop, daemon=True) self.cpu_cache_thread.start() @@ -144,7 +147,7 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes # 超时时,放弃进行 cache page 的匹配。 current_time = time.time() if current_time - start_time >= self.cpu_cache_time_out: - self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + self.send_to_router_queue.put(group_req_indexes) logger.warning( f"cache matching time out {current_time - start_time}s, " f"group_req_id: {group_req_indexes.group_req_id}" @@ -211,14 +214,23 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes for req in reqs: self.shm_req_manager.put_back_req_obj(req) - self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + self.send_to_router_queue.put(group_req_indexes) return + def _send_finished_group_reqs(self): + while True: + try: + group_req_indexes = self.send_to_router_queue.get_nowait() + except Empty: + return + self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + def recv_loop(self): try: recv_max_count = 128 while True: + self._send_finished_group_reqs() recv_objs = [] try: # 一次最多从 zmq 中取 recv_max_count 个请求,防止 zmq 队列中请求数量过多导致阻塞了主循环。 From 94fdfdeca05f3f94619a0c183088114c42e65124 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 27 Aug 2026 22:05:19 +0800 Subject: [PATCH 145/214] perf(dp): move scheduler control to Gloo for decode overlap --- lightllm/distributed/communication_op.py | 3 + .../model_infer/mode_backend/base_backend.py | 47 +++++------ .../mode_backend/dp_backend/control_state.py | 49 ++++-------- .../mode_backend/dp_backend/impl.py | 14 +--- .../mode_backend/test_dp_control.py | 77 +++++++++++++++++++ 5 files changed, 122 insertions(+), 68 deletions(-) create mode 100644 unit_tests/server/router/model_infer/mode_backend/test_dp_control.py diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 306c5c42a8..3fedac78d6 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -157,6 +157,7 @@ def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, as class DistributeGroupManager: def __init__(self): self.groups = [] + self.dp_control_group = None self.ep_balance_monitor_group = None self.ep_buffer = None self.ep_low_latency_buffer = None @@ -176,6 +177,8 @@ def create_groups(self, group_size: int): if not args.disable_flashinfer_allreduce: group.init_flashinfer_reduce() self.groups.append(group) + if args.dp > 1: + self.dp_control_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") if ( getattr(args, "enable_ep_moe", False) and not getattr(args, "disable_ep_balance_monitor", False) 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 0124779aa2..4e990fdd2b 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -211,14 +211,12 @@ def init_model(self, kvargs): ) # 初始化 dp 模式使用的通信 tensor, 对于非dp模式,不会使用到 if self.dp_size > 1: + self.dp_control_tensor = torch.zeros(2, dtype=torch.int32, device="cpu", requires_grad=False) self.dp_reduce_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) - self.dp_gather_item_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) - self.dp_all_gather_tensor = torch.tensor( - [0 for _ in range(self.global_world_size)], dtype=torch.int32, device="cuda", requires_grad=False - ) # 用于协同读取 ShmObjsIOBuffer 中的请求信息的通信tensor和通信组对象。 - self.node_broadcast_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) + self.node_broadcast_tensor = torch.zeros(1, dtype=torch.int32, device="cpu", requires_grad=False) + self.node_gloo_group = create_new_group_for_current_node("gloo") self.node_nccl_group = create_new_group_for_current_node("nccl") # 用于在多节点tp模式下协同读取 ShmObjsIOBuffer 中的请求信息的通信tensor和通信组对象。 @@ -485,8 +483,8 @@ def _try_read_new_reqs_normal(self): self.node_broadcast_tensor.fill_(0) src_rank_id = self.args.node_rank * self.node_world_size - broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_nccl_group, async_op=False) - new_buffer_is_ready = self.node_broadcast_tensor.detach().item() + broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_gloo_group, async_op=False) + new_buffer_is_ready = self.node_broadcast_tensor.item() if new_buffer_is_ready: self._read_reqs_buffer_and_init_reqs() @@ -499,8 +497,8 @@ def _try_read_new_reqs_normal(self): self.node_broadcast_tensor.fill_(0) src_rank_id = self.args.node_rank * self.node_world_size - broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_nccl_group, async_op=False) - new_buffer_is_ready = self.node_broadcast_tensor.detach().item() + broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_gloo_group, async_op=False) + new_buffer_is_ready = self.node_broadcast_tensor.item() if new_buffer_is_ready: self._read_pd_trans_io_buffer_and_update_req_status() return @@ -968,23 +966,26 @@ def _sample_and_scatter_token( ) return next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu - def _dp_all_gather_prefill_and_decode_req_num( + def _dp_all_reduce_req_presence( self, prefill_reqs: List[InferReq], decode_reqs: List[InferReq] - ) -> Tuple[np.ndarray, np.ndarray]: + ) -> tuple[bool, bool]: """ - Gather the number of prefill requests across all DP ranks. - """ - current_dp_prefill_num = len(prefill_reqs) - self.dp_gather_item_tensor.fill_(current_dp_prefill_num) - all_gather_into_tensor(self.dp_all_gather_tensor, self.dp_gather_item_tensor, group=None, async_op=False) - dp_prefill_req_nums = self.dp_all_gather_tensor.cpu().numpy() - - current_dp_decode_num = len(decode_reqs) - self.dp_gather_item_tensor.fill_(current_dp_decode_num) - all_gather_into_tensor(self.dp_all_gather_tensor, self.dp_gather_item_tensor, group=None, async_op=False) - dp_decode_req_nums = self.dp_all_gather_tensor.cpu().numpy() + Return whether any DP rank has prefill or decode requests. - return dp_prefill_req_nums, dp_decode_req_nums + Request counts originate on the CPU and the scheduler only needs their + global presence. Keep this control-plane collective on the CPU so it can + overlap the previous CUDA graph instead of synchronizing that graph back + to the host every decode step. + """ + self.dp_control_tensor[0] = bool(prefill_reqs) + self.dp_control_tensor[1] = bool(decode_reqs) + all_reduce( + self.dp_control_tensor, + op=dist.ReduceOp.MAX, + group=dist_group_manager.dp_control_group, + async_op=False, + ) + return bool(self.dp_control_tensor[0]), bool(self.dp_control_tensor[1]) def _dp_all_reduce_decode_req_num(self, decode_reqs: List[InferReq]) -> int: """ diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py index ffd92202d8..c90a864731 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py @@ -1,8 +1,5 @@ -import numpy as np from enum import Enum -from typing import List from lightllm.utils.envs_utils import get_env_start_args -from lightllm.server.router.model_infer.infer_batch import InferReq from ..base_backend import ModeBackend @@ -20,10 +17,8 @@ def __init__(self, backend: ModeBackend): def select_run_way( self, - dp_prefill_req_nums: np.ndarray, - dp_decode_req_nums: np.ndarray, - prefill_reqs: List[InferReq], - decode_reqs: List[InferReq], + has_prefill: bool, + has_decode: bool, ) -> "RunWay": """ 判断决策运行方式: @@ -32,55 +27,41 @@ def select_run_way( self.step_count += 1 if self.is_aggressive_schedule: return self._agressive_way( - dp_prefill_req_nums=dp_prefill_req_nums, - dp_decode_req_nums=dp_decode_req_nums, - prefill_reqs=prefill_reqs, - decode_reqs=decode_reqs, + has_prefill=has_prefill, + has_decode=has_decode, ) else: return self._normal_way( - dp_prefill_req_nums=dp_prefill_req_nums, - dp_decode_req_nums=dp_decode_req_nums, - prefill_reqs=prefill_reqs, - decode_reqs=decode_reqs, + has_prefill=has_prefill, + has_decode=has_decode, ) def _agressive_way( self, - dp_prefill_req_nums: np.ndarray, - dp_decode_req_nums: np.ndarray, - prefill_reqs: List[InferReq], - decode_reqs: List[InferReq], + has_prefill: bool, + has_decode: bool, ): - max_prefill_num = np.max(dp_prefill_req_nums) - if max_prefill_num > 0: + if has_prefill: return RunWay.PREFILL - max_decode_num = np.max(dp_decode_req_nums) - if max_decode_num > 0: + if has_decode: return RunWay.DECODE return RunWay.PASS def _normal_way( self, - dp_prefill_req_nums: np.ndarray, - dp_decode_req_nums: np.ndarray, - prefill_reqs: List[InferReq], - decode_reqs: List[InferReq], + has_prefill: bool, + has_decode: bool, ): - # use_ratio = np.count_nonzero(dp_prefill_req_nums) / dp_prefill_req_nums.shape[0] - max_decode_num = np.max(dp_decode_req_nums) - max_prefill_num = np.max(dp_prefill_req_nums) - - if self.left_decode_num > 0 and max_decode_num > 0: + if self.left_decode_num > 0 and has_decode: self.left_decode_num -= 1 return RunWay.DECODE - if max_prefill_num > 0: + if has_prefill: # prefill 一次允许进行几次 decode 操作。 self.left_decode_num = self.decode_max_step return RunWay.PREFILL else: - if max_decode_num > 0: + if has_decode: return RunWay.DECODE else: return RunWay.PASS 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 d9a28bdbfc..c712d87ecb 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 @@ -136,21 +136,13 @@ def infer_loop(self): recover_paused=self.control_state_machine.try_recover_paused_reqs(), ) - if self.support_overlap: - # The previous infer thread releases forward before its GPU work finishes. - # Keep the next DP collective behind that work to avoid overlapping NCCL - # with the previous DeepEP/MTP kernels. - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - - dp_prefill_req_nums, dp_decode_req_nums = self._dp_all_gather_prefill_and_decode_req_num( + has_prefill, has_decode = self._dp_all_reduce_req_presence( prefill_reqs=prefill_reqs, decode_reqs=decode_reqs ) run_way = self.control_state_machine.select_run_way( - dp_prefill_req_nums=dp_prefill_req_nums, - dp_decode_req_nums=dp_decode_req_nums, - prefill_reqs=prefill_reqs, - decode_reqs=decode_reqs, + has_prefill=has_prefill, + has_decode=has_decode, ) if run_way.is_prefill(): diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py new file mode 100644 index 0000000000..51358bd1f8 --- /dev/null +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py @@ -0,0 +1,77 @@ +from types import SimpleNamespace + +import torch + +from lightllm.server.router.model_infer.mode_backend import base_backend as base_backend_module +from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend +from lightllm.server.router.model_infer.mode_backend.dp_backend.control_state import DPControlState, RunWay + + +def test_dp_req_presence_uses_one_cpu_collective(monkeypatch): + control_group = object() + backend = ModeBackend.__new__(ModeBackend) + backend.dp_control_tensor = torch.zeros(2, dtype=torch.int32, device="cpu") + calls = [] + + def all_reduce(tensor, op, group, async_op): + calls.append((tensor.tolist(), tensor.device.type, op, group, async_op)) + tensor.fill_(1) + + monkeypatch.setattr(base_backend_module.dist_group_manager, "dp_control_group", control_group) + monkeypatch.setattr(base_backend_module, "all_reduce", all_reduce) + + has_prefill, has_decode = backend._dp_all_reduce_req_presence(prefill_reqs=[object()], decode_reqs=[]) + + assert (has_prefill, has_decode) == (True, True) + assert calls == [ + ([1, 0], "cpu", torch.distributed.ReduceOp.MAX, control_group, False), + ] + + +def test_new_request_readiness_uses_cpu_control_group(monkeypatch): + control_group = object() + backend = ModeBackend.__new__(ModeBackend) + backend.is_master_in_node = True + backend.shm_reqs_io_buffer = SimpleNamespace(is_ready=lambda: True) + backend.node_broadcast_tensor = torch.zeros(1, dtype=torch.int32, device="cpu") + backend.node_gloo_group = control_group + backend.node_world_size = 8 + backend.args = SimpleNamespace(node_rank=0) + backend.is_pd_mode = False + calls = [] + + def broadcast(tensor, src, group, async_op): + calls.append((tensor.tolist(), tensor.device.type, src, group, async_op)) + + backend._read_reqs_buffer_and_init_reqs = lambda: calls.append("read") + monkeypatch.setattr(base_backend_module, "broadcast", broadcast) + + backend._try_read_new_reqs_normal() + + assert calls == [ + ([1], "cpu", 0, control_group, False), + "read", + ] + + +def test_aggressive_dp_control_prefers_prefill(): + state = DPControlState.__new__(DPControlState) + state.is_aggressive_schedule = True + state.step_count = 0 + + assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.PREFILL + assert state.select_run_way(has_prefill=False, has_decode=True) is RunWay.DECODE + assert state.select_run_way(has_prefill=False, has_decode=False) is RunWay.PASS + assert state.step_count == 3 + + +def test_normal_dp_control_preserves_decode_budget(): + state = DPControlState.__new__(DPControlState) + state.is_aggressive_schedule = False + state.decode_max_step = 1 + state.left_decode_num = 1 + state.step_count = 0 + + assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.DECODE + assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.PREFILL + assert state.left_decode_num == 1 From b70bc7ac0d5ee43e63c48520914642ffe8ac4ac8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 27 Aug 2026 22:05:28 +0800 Subject: [PATCH 146/214] perf(dsv4): add five-stream decode overlap --- .../layer_infer/transformer_layer_infer.py | 220 +++++++++++++++--- lightllm/models/deepseek_v4/model.py | 4 + 2 files changed, 198 insertions(+), 26 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index ecb147e5a2..4b123b0ac3 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,3 +1,5 @@ +from contextlib import nullcontext + import torch import triton from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl @@ -20,6 +22,7 @@ _C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 +_DECODE_MULTI_STREAM_MAX_BATCH_SIZE = 64 class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): @@ -61,6 +64,14 @@ def __init__(self, layer_num, network_config): layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ ) self.dsv4_prefill_aux_stream = None + self.dsv4_decode_aux_streams = None + + def _use_decode_multi_stream(self, infer_state: DeepseekV4InferStateInfo, token_num: int) -> bool: + return ( + self.dsv4_decode_aux_streams is not None + and infer_state.is_cuda_graph + and token_num <= _DECODE_MULTI_STREAM_MAX_BATCH_SIZE + ) # ------------------------------------------------------------------ forward (HC-threaded) def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): @@ -487,16 +498,104 @@ def _context_attention_kernel( def token_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): + if self._use_decode_multi_stream(infer_state, x.shape[0]): + return self._token_attention_forward_multi_stream(x, infer_state, layer_weight) q, q_lora, full_x = self._get_qkv(x, infer_state, layer_weight) o = self._token_attention_kernel(q, q_lora, full_x, infer_state, layer_weight) return self._get_o(o, infer_state, layer_weight) + def _token_attention_forward_multi_stream( + self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + ): + from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import fused_q_norm_rope + + full_x = self._tpsp_allgather(input=x, infer_state=infer_state) + token_num = full_x.shape[0] + main_stream = torch.cuda.current_stream() + kv_stream, compressor_stream, indexer_stream = self.dsv4_decode_aux_streams[:3] + indexer_aux_streams = self.dsv4_decode_aux_streams[3:] + + if self.compress_ratio: + compressor_stream.wait_stream(main_stream) + with torch.cuda.stream(compressor_stream): + full_x.record_stream(compressor_stream) + self.compressor.compress( + full_x, + infer_state, + layer_weight, + self.cos_compress_table, + self.sin_compress_table, + use_custom_tensor_manager=False, + ) + + qkv = layer_weight.wq_a_wkv_.mm(full_x) + kv_stream.wait_stream(main_stream) + with torch.cuda.stream(kv_stream): + qkv.record_stream(kv_stream) + infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( + layer_index=self.layer_num_, + mem_index=infer_state.mem_index, + swa_slots=getattr(infer_state, "dsv4_swa_write_slots", None), + kv=qkv[:, -self.head_dim_ :], + kv_weight=layer_weight.kv_norm_.weight, + eps=self.eps_, + freqs_cis=self.freqs_cis, + positions=infer_state.position_ids, + ) + + q_lora = layer_weight.q_norm_(qkv[:, : -self.head_dim_], eps=self.eps_) + if self.compress_ratio == 4: + indexer_stream.wait_stream(main_stream) + with torch.cuda.stream(indexer_stream): + full_x.record_stream(indexer_stream) + q_lora.record_stream(indexer_stream) + self.index_infer.write_indexer_k( + full_x, + infer_state, + layer_weight, + self.cos_compress_table, + self.sin_compress_table, + use_custom_tensor_manager=False, + ) + meta = self.index_infer.build_metadata( + full_x, + q_lora, + infer_state, + layer_weight, + use_custom_tensor_manager=False, + aux_streams=indexer_aux_streams, + ) + else: + meta = self.index_infer.build_metadata(full_x, q_lora, infer_state, layer_weight) + + q_in = layer_weight.wq_b_.mm(q_lora).view(token_num, self.tp_q_head_num_, self.head_dim_) + q = self.alloc_tensor( + (token_num, self.flashmla_q_head_num_, self.head_dim_), dtype=q_in.dtype, device=q_in.device + ) + q[:, self.tp_q_head_num_ :, :].zero_() + fused_q_norm_rope(q_in, q[:, : self.tp_q_head_num_, :], self.eps_, self.freqs_cis, infer_state.position_ids) + + main_stream.wait_stream(kv_stream) + if self.compress_ratio: + main_stream.wait_stream(compressor_stream) + if self.compress_ratio == 4: + main_stream.wait_stream(indexer_stream) + for tensor in (meta.get("extra_indices"), meta.get("extra_lengths")): + if tensor is not None: + tensor.record_stream(main_stream) + + o = self._run_token_attention(q, meta, infer_state, layer_weight) + return self._get_o(o, infer_state, layer_weight) + def _token_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): self.compressor.compress(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) + return self._run_token_attention(q, meta, infer_state, layer_weight) + + def _run_token_attention(self, q, meta, infer_state, layer_weight): att_control = AttControl( nsa_decode=True, nsa_decode_dict={ @@ -534,14 +633,21 @@ def _routed_experts( alloc_tensor_func=self.alloc_tensor, ) - def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): + def _ffn_tp( + self, + input, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + use_custom_tensor_manager: bool = True, + ): input = input.view(-1, self.embed_dim_) - gate_up = layer_weight.gate_up_proj.mm(input) - shared = self.alloc_tensor((input.size(0), gate_up.size(1) // 2), input.dtype) + gate_up = layer_weight.gate_up_proj.mm(input, use_custom_tensor_mananger=use_custom_tensor_manager) + alloc_tensor = self.alloc_tensor if use_custom_tensor_manager else torch.empty + shared = alloc_tensor((input.size(0), gate_up.size(1) // 2), dtype=input.dtype, device=input.device) silu_and_mul_fwd(gate_up, shared, limit=self.swiglu_limit) input = None gate_up = None - out = layer_weight.down_proj.mm(shared) + out = layer_weight.down_proj.mm(shared, use_custom_tensor_mananger=use_custom_tensor_manager) shared = None return out @@ -550,14 +656,31 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV x = self._tpsp_allgather(input=x, infer_state=infer_state) logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) - weights, indices = self._select_experts(logits, infer_state, layer_weight) - # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts - # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 + # The non-EP FuseMoeTriton path can overwrite x, so it keeps shared first. + # EP uses the read-only DeepEP/DeepGEMM path and may overlap both branches. # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. - shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) + use_multi_stream = self.enable_ep_moe and self._use_decode_multi_stream(infer_state, x.shape[0]) + if use_multi_stream: + main_stream = torch.cuda.current_stream() + shared_stream = self.dsv4_decode_aux_streams[0] + shared_stream.wait_stream(main_stream) + with torch.cuda.stream(shared_stream): + x.record_stream(shared_stream) + shared = self._ffn_tp( + input=x, + infer_state=infer_state, + layer_weight=layer_weight, + use_custom_tensor_manager=False, + ) + weights, indices = self._select_experts(logits, infer_state, layer_weight) + if not use_multi_stream: + shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) if self.enable_ep_moe: + if use_multi_stream: + main_stream.wait_stream(shared_stream) + shared.record_stream(main_stream) return routed + shared out = routed + shared return self._tpsp_reduce(input=out, infer_state=infer_state) @@ -735,7 +858,13 @@ def write_indexer_k( ) def build_metadata( - self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + self, + x, + q_lora, + infer_state: DeepseekV4InferStateInfo, + layer_weight, + use_custom_tensor_manager=True, + aux_streams=None, ): swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) swa_lengths = infer_state.dsv4_swa_lengths @@ -743,7 +872,12 @@ def build_metadata( extra_indices = extra_lengths = None if self.compress_ratio == 4: idx_q_fp8, weights = self._indexer_q_weight( - x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager + x, + q_lora, + infer_state, + layer_weight, + use_custom_tensor_manager=use_custom_tensor_manager, + aux_streams=aux_streams, ) extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: @@ -757,7 +891,13 @@ def build_metadata( } def _indexer_q_weight( - self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + self, + x, + q_lora, + infer_state: DeepseekV4InferStateInfo, + layer_weight, + use_custom_tensor_manager=True, + aux_streams=None, ): # Fused: wq_b mm -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant, with the # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: @@ -772,25 +912,53 @@ def _indexer_q_weight( raise RuntimeError( f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" ) - idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( - token_num, self.index_n_heads, self.index_head_dim - ) - raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( - token_num, self.index_n_heads - ) # [T, H] raw + if aux_streams is None: + idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads, self.index_head_dim + ) + raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads + ) + else: + assert not use_custom_tensor_manager + q_stream, weight_stream = aux_streams + current_stream = torch.cuda.current_stream() + q_stream.wait_stream(current_stream) + weight_stream.wait_stream(current_stream) + with torch.cuda.stream(q_stream): + q_lora.record_stream(q_stream) + idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=False).view( + token_num, self.index_n_heads, self.index_head_dim + ) + with torch.cuda.stream(weight_stream): + x.record_stream(weight_stream) + raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=False).view( + token_num, self.index_n_heads + ) + q_stream.wait_stream(weight_stream) + idx_q_fp8_out = weights_out = None if use_custom_tensor_manager: idx_q_fp8_out = self.alloc_tensor(idx_q.shape, torch.float8_e4m3fn, device=idx_q.device) weights_out = self.alloc_tensor((*idx_q.shape[:-1], 1), torch.float32, device=idx_q.device) - idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( - idx_q, - raw_w, - self.indexer_weight_scale, - self.freqs_cis, - infer_state.position_ids, - q_fp8=idx_q_fp8_out, - weights_out=weights_out, - ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded + target_stream = q_stream if aux_streams is not None else None + stream_context = torch.cuda.stream(target_stream) if target_stream is not None else nullcontext() + with stream_context: + if target_stream is not None: + raw_w.record_stream(target_stream) + idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( + idx_q, + raw_w, + self.indexer_weight_scale, + self.freqs_cis, + infer_state.position_ids, + q_fp8=idx_q_fp8_out, + weights_out=weights_out, + ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded + if target_stream is not None: + current_stream.wait_stream(target_stream) + idx_q_fp8.record_stream(current_stream) + weights.record_stream(current_stream) return idx_q_fp8, weights.squeeze(-1) def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 5f236d2e4b..bd07f680bc 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -154,6 +154,10 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() self.dsv4_workspace = DeepseekV4Workspace(self) + if self.args.enable_ep_moe and not self.args.enable_decode_microbatch_overlap: + decode_aux_streams = tuple(torch.cuda.Stream() for _ in range(5)) + for layer in self.layers_infer: + layer.dsv4_decode_aux_streams = decode_aux_streams if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1" and not self.args.enable_prefill_microbatch_overlap: prefill_aux_stream = torch.cuda.Stream() for layer in self.layers_infer: From 4f60ca53236281bb070609b555cb1f5c890c54d2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 27 Aug 2026 22:05:37 +0800 Subject: [PATCH 147/214] perf(dsv4): eliminate CUDA list-index sync in decode slot prep --- .../deepseek4_mem_manager.py | 12 +++++++++--- lightllm/common/req_manager.py | 17 +++++++++++++---- 2 files changed, 22 insertions(+), 7 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index c2c5b787f3..da12d29b8f 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -889,15 +889,21 @@ def alloc_swa_decode( cont_rows.append(i) mem_indexes = mem_indexes.reshape(-1) if cont_rows: - prev_full = prev_full_indexes.reshape(-1)[cont_rows] + # Steady decode normally puts every request on the same page offset. + # Avoid Python-list indexing in that case: PyTorch copies the list to + # CUDA synchronously and waits for the previous decode graph. + all_rows = len(cont_rows) == len(req_list) + prev_full = prev_full_indexes.reshape(-1) if all_rows else prev_full_indexes.reshape(-1)[cont_rows] prev_slots = self.full_to_swa_indexs[prev_full] slots = prev_slots + 1 - self.full_to_swa_indexs[mem_indexes[cont_rows]] = slots + dst_indexes = mem_indexes if all_rows else mem_indexes[cont_rows] + self.full_to_swa_indexs[dst_indexes] = slots self._update_swa_page_counts(slots, 1) if new_rows: pages = self._alloc_swa_pages(len(new_rows)).to(self.full_to_swa_indexs.device, non_blocking=True) slots = pages * page - self.full_to_swa_indexs[mem_indexes[new_rows]] = slots + dst_indexes = mem_indexes if len(new_rows) == len(req_list) else mem_indexes[new_rows] + self.full_to_swa_indexs[dst_indexes] = slots self._update_swa_page_counts(slots, 1) return diff --git a/lightllm/common/req_manager.py b/lightllm/common/req_manager.py index 31603d84ff..4c3eec7213 100644 --- a/lightllm/common/req_manager.py +++ b/lightllm/common/req_manager.py @@ -759,14 +759,20 @@ def _scatter_c4_decode_slots( prev_meta = prev_meta.view(-1, 2) prev_full = self.req_to_token_indexs[prev_meta[:, 0], prev_meta[:, 1]] else: - prev_full = prev_group_end_mem_indexes.reshape(-1)[cont_rows] + prev_full = ( + prev_group_end_mem_indexes.reshape(-1) + if len(cont_rows) == len(req_list) + else prev_group_end_mem_indexes.reshape(-1)[cont_rows] + ) prev_slots = mapping[prev_full] - self._register_c4_slots(mem_indexes[cont_rows], prev_slots + 1) + dst_indexes = mem_indexes if len(cont_rows) == len(req_list) else mem_indexes[cont_rows] + self._register_c4_slots(dst_indexes, prev_slots + 1) if new_rows: self._realize_c4_pages(len(new_rows)) # 兑现: 精确需求, 复用已算的 new_rows pages = self.mem_manager.alloc_c4_pages(len(new_rows)).to(mapping.device, non_blocking=True) - self._register_c4_slots(mem_indexes[new_rows], pages * page) + dst_indexes = mem_indexes if len(new_rows) == len(req_list) else mem_indexes[new_rows] + self._register_c4_slots(dst_indexes, pages * page) return def _scatter_c128_slots(self, full_slots: torch.Tensor) -> None: @@ -866,7 +872,10 @@ def prepare_decode_compress_slots( if req_idx != self.HOLD_REQUEST_ID and seq_len > 0 and seq_len % ratio == 0 ] if rows: - self._scatter_c128_slots(mem_indexes.reshape(-1)[rows]) + full_slots = mem_indexes.reshape(-1) + if len(rows) != len(req_list): + full_slots = full_slots[rows] + self._scatter_c128_slots(full_slots) return def alloc(self): From 8991e230d9a96825857baf1bfbe1b922b3cd7eb7 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 28 Aug 2026 03:31:55 +0000 Subject: [PATCH 148/214] perf(dsv4): remove decode multi-stream overlap --- .../layer_infer/transformer_layer_infer.py | 220 +++--------------- lightllm/models/deepseek_v4/model.py | 4 - 2 files changed, 26 insertions(+), 198 deletions(-) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 4b123b0ac3..ecb147e5a2 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -1,5 +1,3 @@ -from contextlib import nullcontext - import torch import triton from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl @@ -22,7 +20,6 @@ _C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 -_DECODE_MULTI_STREAM_MAX_BATCH_SIZE = 64 class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): @@ -64,14 +61,6 @@ def __init__(self, layer_num, network_config): layer_idx=self.layer_num_, network_config=self.network_config_, tp_world_size=self.tp_world_size_ ) self.dsv4_prefill_aux_stream = None - self.dsv4_decode_aux_streams = None - - def _use_decode_multi_stream(self, infer_state: DeepseekV4InferStateInfo, token_num: int) -> bool: - return ( - self.dsv4_decode_aux_streams is not None - and infer_state.is_cuda_graph - and token_num <= _DECODE_MULTI_STREAM_MAX_BATCH_SIZE - ) # ------------------------------------------------------------------ forward (HC-threaded) def _hc_attn_in(self, input_embdings, layer_weight: DeepseekV4TransformerLayerWeight): @@ -498,104 +487,16 @@ def _context_attention_kernel( def token_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): - if self._use_decode_multi_stream(infer_state, x.shape[0]): - return self._token_attention_forward_multi_stream(x, infer_state, layer_weight) q, q_lora, full_x = self._get_qkv(x, infer_state, layer_weight) o = self._token_attention_kernel(q, q_lora, full_x, infer_state, layer_weight) return self._get_o(o, infer_state, layer_weight) - def _token_attention_forward_multi_stream( - self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight - ): - from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import fused_q_norm_rope - - full_x = self._tpsp_allgather(input=x, infer_state=infer_state) - token_num = full_x.shape[0] - main_stream = torch.cuda.current_stream() - kv_stream, compressor_stream, indexer_stream = self.dsv4_decode_aux_streams[:3] - indexer_aux_streams = self.dsv4_decode_aux_streams[3:] - - if self.compress_ratio: - compressor_stream.wait_stream(main_stream) - with torch.cuda.stream(compressor_stream): - full_x.record_stream(compressor_stream) - self.compressor.compress( - full_x, - infer_state, - layer_weight, - self.cos_compress_table, - self.sin_compress_table, - use_custom_tensor_manager=False, - ) - - qkv = layer_weight.wq_a_wkv_.mm(full_x) - kv_stream.wait_stream(main_stream) - with torch.cuda.stream(kv_stream): - qkv.record_stream(kv_stream) - infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( - layer_index=self.layer_num_, - mem_index=infer_state.mem_index, - swa_slots=getattr(infer_state, "dsv4_swa_write_slots", None), - kv=qkv[:, -self.head_dim_ :], - kv_weight=layer_weight.kv_norm_.weight, - eps=self.eps_, - freqs_cis=self.freqs_cis, - positions=infer_state.position_ids, - ) - - q_lora = layer_weight.q_norm_(qkv[:, : -self.head_dim_], eps=self.eps_) - if self.compress_ratio == 4: - indexer_stream.wait_stream(main_stream) - with torch.cuda.stream(indexer_stream): - full_x.record_stream(indexer_stream) - q_lora.record_stream(indexer_stream) - self.index_infer.write_indexer_k( - full_x, - infer_state, - layer_weight, - self.cos_compress_table, - self.sin_compress_table, - use_custom_tensor_manager=False, - ) - meta = self.index_infer.build_metadata( - full_x, - q_lora, - infer_state, - layer_weight, - use_custom_tensor_manager=False, - aux_streams=indexer_aux_streams, - ) - else: - meta = self.index_infer.build_metadata(full_x, q_lora, infer_state, layer_weight) - - q_in = layer_weight.wq_b_.mm(q_lora).view(token_num, self.tp_q_head_num_, self.head_dim_) - q = self.alloc_tensor( - (token_num, self.flashmla_q_head_num_, self.head_dim_), dtype=q_in.dtype, device=q_in.device - ) - q[:, self.tp_q_head_num_ :, :].zero_() - fused_q_norm_rope(q_in, q[:, : self.tp_q_head_num_, :], self.eps_, self.freqs_cis, infer_state.position_ids) - - main_stream.wait_stream(kv_stream) - if self.compress_ratio: - main_stream.wait_stream(compressor_stream) - if self.compress_ratio == 4: - main_stream.wait_stream(indexer_stream) - for tensor in (meta.get("extra_indices"), meta.get("extra_lengths")): - if tensor is not None: - tensor.record_stream(main_stream) - - o = self._run_token_attention(q, meta, infer_state, layer_weight) - return self._get_o(o, infer_state, layer_weight) - def _token_attention_kernel( self, q, q_lora, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): self.compressor.compress(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) self.index_infer.write_indexer_k(x, infer_state, layer_weight, self.cos_compress_table, self.sin_compress_table) meta = self.index_infer.build_metadata(x, q_lora, infer_state, layer_weight) - return self._run_token_attention(q, meta, infer_state, layer_weight) - - def _run_token_attention(self, q, meta, infer_state, layer_weight): att_control = AttControl( nsa_decode=True, nsa_decode_dict={ @@ -633,21 +534,14 @@ def _routed_experts( alloc_tensor_func=self.alloc_tensor, ) - def _ffn_tp( - self, - input, - infer_state: DeepseekV4InferStateInfo, - layer_weight: DeepseekV4TransformerLayerWeight, - use_custom_tensor_manager: bool = True, - ): + def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): input = input.view(-1, self.embed_dim_) - gate_up = layer_weight.gate_up_proj.mm(input, use_custom_tensor_mananger=use_custom_tensor_manager) - alloc_tensor = self.alloc_tensor if use_custom_tensor_manager else torch.empty - shared = alloc_tensor((input.size(0), gate_up.size(1) // 2), dtype=input.dtype, device=input.device) + gate_up = layer_weight.gate_up_proj.mm(input) + shared = self.alloc_tensor((input.size(0), gate_up.size(1) // 2), input.dtype) silu_and_mul_fwd(gate_up, shared, limit=self.swiglu_limit) input = None gate_up = None - out = layer_weight.down_proj.mm(shared, use_custom_tensor_mananger=use_custom_tensor_manager) + out = layer_weight.down_proj.mm(shared) shared = None return out @@ -656,31 +550,14 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV x = self._tpsp_allgather(input=x, infer_state=infer_state) logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) - # The non-EP FuseMoeTriton path can overwrite x, so it keeps shared first. - # EP uses the read-only DeepEP/DeepGEMM path and may overlap both branches. + weights, indices = self._select_experts(logits, infer_state, layer_weight) + # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts + # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. - use_multi_stream = self.enable_ep_moe and self._use_decode_multi_stream(infer_state, x.shape[0]) - if use_multi_stream: - main_stream = torch.cuda.current_stream() - shared_stream = self.dsv4_decode_aux_streams[0] - shared_stream.wait_stream(main_stream) - with torch.cuda.stream(shared_stream): - x.record_stream(shared_stream) - shared = self._ffn_tp( - input=x, - infer_state=infer_state, - layer_weight=layer_weight, - use_custom_tensor_manager=False, - ) - weights, indices = self._select_experts(logits, infer_state, layer_weight) - if not use_multi_stream: - shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) + shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) if self.enable_ep_moe: - if use_multi_stream: - main_stream.wait_stream(shared_stream) - shared.record_stream(main_stream) return routed + shared out = routed + shared return self._tpsp_reduce(input=out, infer_state=infer_state) @@ -858,13 +735,7 @@ def write_indexer_k( ) def build_metadata( - self, - x, - q_lora, - infer_state: DeepseekV4InferStateInfo, - layer_weight, - use_custom_tensor_manager=True, - aux_streams=None, + self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True ): swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) swa_lengths = infer_state.dsv4_swa_lengths @@ -872,12 +743,7 @@ def build_metadata( extra_indices = extra_lengths = None if self.compress_ratio == 4: idx_q_fp8, weights = self._indexer_q_weight( - x, - q_lora, - infer_state, - layer_weight, - use_custom_tensor_manager=use_custom_tensor_manager, - aux_streams=aux_streams, + x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager ) extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: @@ -891,13 +757,7 @@ def build_metadata( } def _indexer_q_weight( - self, - x, - q_lora, - infer_state: DeepseekV4InferStateInfo, - layer_weight, - use_custom_tensor_manager=True, - aux_streams=None, + self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True ): # Fused: wq_b mm -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant, with the # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: @@ -912,53 +772,25 @@ def _indexer_q_weight( raise RuntimeError( f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" ) - if aux_streams is None: - idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( - token_num, self.index_n_heads, self.index_head_dim - ) - raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( - token_num, self.index_n_heads - ) - else: - assert not use_custom_tensor_manager - q_stream, weight_stream = aux_streams - current_stream = torch.cuda.current_stream() - q_stream.wait_stream(current_stream) - weight_stream.wait_stream(current_stream) - with torch.cuda.stream(q_stream): - q_lora.record_stream(q_stream) - idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=False).view( - token_num, self.index_n_heads, self.index_head_dim - ) - with torch.cuda.stream(weight_stream): - x.record_stream(weight_stream) - raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=False).view( - token_num, self.index_n_heads - ) - q_stream.wait_stream(weight_stream) - + idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads, self.index_head_dim + ) + raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( + token_num, self.index_n_heads + ) # [T, H] raw idx_q_fp8_out = weights_out = None if use_custom_tensor_manager: idx_q_fp8_out = self.alloc_tensor(idx_q.shape, torch.float8_e4m3fn, device=idx_q.device) weights_out = self.alloc_tensor((*idx_q.shape[:-1], 1), torch.float32, device=idx_q.device) - target_stream = q_stream if aux_streams is not None else None - stream_context = torch.cuda.stream(target_stream) if target_stream is not None else nullcontext() - with stream_context: - if target_stream is not None: - raw_w.record_stream(target_stream) - idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( - idx_q, - raw_w, - self.indexer_weight_scale, - self.freqs_cis, - infer_state.position_ids, - q_fp8=idx_q_fp8_out, - weights_out=weights_out, - ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded - if target_stream is not None: - current_stream.wait_stream(target_stream) - idx_q_fp8.record_stream(current_stream) - weights.record_stream(current_stream) + idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( + idx_q, + raw_w, + self.indexer_weight_scale, + self.freqs_cis, + infer_state.position_ids, + q_fp8=idx_q_fp8_out, + weights_out=weights_out, + ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded return idx_q_fp8, weights.squeeze(-1) def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index bd07f680bc..5f236d2e4b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -154,10 +154,6 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() self.dsv4_workspace = DeepseekV4Workspace(self) - if self.args.enable_ep_moe and not self.args.enable_decode_microbatch_overlap: - decode_aux_streams = tuple(torch.cuda.Stream() for _ in range(5)) - for layer in self.layers_infer: - layer.dsv4_decode_aux_streams = decode_aux_streams if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1" and not self.args.enable_prefill_microbatch_overlap: prefill_aux_stream = torch.cuda.Stream() for layer in self.layers_infer: From c908c766e791c40904ea798237ed62949f08b660 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 28 Aug 2026 05:07:00 +0000 Subject: [PATCH 149/214] format --- .../triton_kernel/build_dspark_swa_index.py | 4 +--- .../models/deepseek_v4_dspark/infer_struct.py | 6 +----- .../layer_infer/post_layer_infer.py | 16 ++++------------ .../layer_infer/transformer_layer_infer.py | 4 +--- .../layer_weights/pre_and_post_layer_weight.py | 6 +++--- lightllm/models/deepseek_v4_dspark/model.py | 3 +-- .../router/model_infer/mtp_speculative/utils.py | 4 +--- 7 files changed, 12 insertions(+), 31 deletions(-) diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py index a320a15068..9b841159ac 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py @@ -62,9 +62,7 @@ def _build_dspark_swa_index_kernel( valid = is_history | is_block output = tl.where(valid, swa_slot, -1) output = tl.where(is_hold, HOLD_SWA_SLOT, output).to(tl.int32) - tl.store( - swa_index_ptr + token_idx * swa_index_stride0 + column, output, mask=in_width - ) + tl.store(swa_index_ptr + token_idx * swa_index_stride0 + column, output, mask=in_width) length = tl.where(is_hold, 1, history_len + BLOCK_SIZE).to(tl.int32) tl.store(swa_length_ptr + token_idx, length) diff --git a/lightllm/models/deepseek_v4_dspark/infer_struct.py b/lightllm/models/deepseek_v4_dspark/infer_struct.py index 73cb48c220..14837466f8 100644 --- a/lightllm/models/deepseek_v4_dspark/infer_struct.py +++ b/lightllm/models/deepseek_v4_dspark/infer_struct.py @@ -18,11 +18,7 @@ def init_some_extra_state(self, model): # pages and replace these base indices below. return - ( - self.dsv4_swa_indices, - self.dsv4_swa_lengths, - self.dsv4_swa_write_slots, - ) = model.dsv4_workspace.dspark_swa( + (self.dsv4_swa_indices, self.dsv4_swa_lengths, self.dsv4_swa_write_slots,) = model.dsv4_workspace.dspark_swa( self.microbatch_index, self.position_ids.numel(), ) diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py index 002ca26fdd..81aa1a269b 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py @@ -25,12 +25,8 @@ def predict_confidence_logits( dim=1, ) prev_embeddings = self._markov_prev_embeddings(prev_token_ids, layer_weight) - features = torch.cat( - [block_hidden, prev_embeddings.to(dtype=block_hidden.dtype)], dim=-1 - ) - logits = layer_weight.confidence_head_weight_.mm( - features.flatten(0, -2).float() - ) + features = torch.cat([block_hidden, prev_embeddings.to(dtype=block_hidden.dtype)], dim=-1) + logits = layer_weight.confidence_head_weight_.mm(features.flatten(0, -2).float()) return logits.view(features.shape[:-1]) def token_forward( @@ -60,15 +56,11 @@ def token_forward( last_input, token_num = self._slice_get_last_input(collapsed, 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 - ] + anchor_token_ids = infer_state.input_ids.reshape(num_reqs, self.block_size_)[:, 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 - ) + 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, diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py index 9006434671..72d08bb720 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py @@ -16,9 +16,7 @@ def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) final_layer = network_config["n_layer"] + network_config["dspark_layer_num"] - 1 self.is_last_layer = layer_num == final_layer - assert ( - self.compress_ratio == 0 - ), "DeepSeek-V4 DSpark draft layers must be SWA-only" + assert self.compress_ratio == 0, "DeepSeek-V4 DSpark draft layers must be SWA-only" def context_forward( self, diff --git a/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py b/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py index f5159f563e..22ae2bb5a4 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py +++ b/lightllm/models/deepseek_v4_dspark/layer_weights/pre_and_post_layer_weight.py @@ -122,9 +122,9 @@ def _dequant_main_proj_in_place(self, weights): scale_name = self.main_proj_weight_.weight_scale_names[0] if scale_name is None: - weights[weight_name] = dequant_fp8_block_to_bf16( - weights[weight_name], weights[scale_key] - ).to(self.data_type_) + weights[weight_name] = dequant_fp8_block_to_bf16(weights[weight_name], weights[scale_key]).to( + self.data_type_ + ) else: weights[scale_name] = weights[scale_key].to(torch.float32) del weights[scale_key] diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index f483b54645..458f162b88 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -79,8 +79,7 @@ def _verify_params(self): assert not self.enable_tpsp_mix_mode, "DeepSeek-V4 DSpark draft model does not support TP-SP" checkpoint_block_size = int(self.config["dspark_block_size"]) assert 0 < self.args.mtp_step <= checkpoint_block_size, ( - f"DeepSeek-V4 DSpark requires --mtp_step in [1, {checkpoint_block_size}], " - f"got {self.args.mtp_step}" + f"DeepSeek-V4 DSpark requires --mtp_step in [1, {checkpoint_block_size}], " f"got {self.args.mtp_step}" ) # The checkpoint block size is its maximum supported width. Runtime # allocation, indexing, and Markov decoding follow --mtp_step. diff --git a/lightllm/server/router/model_infer/mtp_speculative/utils.py b/lightllm/server/router/model_infer/mtp_speculative/utils.py index de4758ebee..38e4ebe651 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/utils.py +++ b/lightllm/server/router/model_infer/mtp_speculative/utils.py @@ -129,9 +129,7 @@ def free_mem_indexes( mem_indexes_to_free.append(extra_indexes_cpu) else: assert extra_mem_to_free.free_mask_cpu is None - dspark_scratch_to_free.append( - (extra_indexes_cpu, extra_mem_to_free.swa_pages_cpu) - ) + dspark_scratch_to_free.append((extra_indexes_cpu, extra_mem_to_free.swa_pages_cpu)) mem_manager = backend.model.req_manager.mem_manager for mem_indexes_cpu, pages_cpu in dspark_scratch_to_free: From 84232e97d0ae540bf72d6dbbd7d9016a2879d79a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 28 Aug 2026 05:17:51 +0000 Subject: [PATCH 150/214] format --- test/unit/test_deepseek_v4_dspark.py | 12 +- tools/convert_deepseek_v4_mtp_to_bf16.py | 381 ++++++++++++++--------- 2 files changed, 236 insertions(+), 157 deletions(-) diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index f972e46a6e..e4432852b2 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -96,15 +96,11 @@ def test_dspark_cuda_graph_padding_extends_only_gpu_scratch_pages(): mtp_draft_swa_pages=pages_cpu.clone(), ) - padded_input = model._create_padded_decode_model_input( - model_input, graph_batch_size - ) + padded_input = model._create_padded_decode_model_input(model_input, graph_batch_size) assert padded_input.mtp_draft_swa_pages.shape == (16,) torch.testing.assert_close(padded_input.mtp_draft_swa_pages[:14], pages_cpu) - torch.testing.assert_close( - padded_input.mtp_draft_swa_pages[14:], torch.zeros(2, dtype=torch.int32) - ) + torch.testing.assert_close(padded_input.mtp_draft_swa_pages[14:], torch.zeros(2, dtype=torch.int32)) assert padded_input.mtp_draft_swa_pages_cpu is pages_cpu @@ -242,9 +238,7 @@ def test_build_dspark_swa_index_exposes_history_and_complete_block(): device="cuda", ) torch.testing.assert_close(indices, expected_indices) - torch.testing.assert_close( - lengths, torch.tensor([7, 7, 7, 5, 5, 5], dtype=torch.int32, device="cuda") - ) + torch.testing.assert_close(lengths, torch.tensor([7, 7, 7, 5, 5, 5], dtype=torch.int32, device="cuda")) torch.testing.assert_close( write_slots, torch.tensor([256, 257, 258, 384, 385, 386], dtype=torch.int32, device="cuda"), diff --git a/tools/convert_deepseek_v4_mtp_to_bf16.py b/tools/convert_deepseek_v4_mtp_to_bf16.py index d0a2cfcb5b..ca58304e2d 100644 --- a/tools/convert_deepseek_v4_mtp_to_bf16.py +++ b/tools/convert_deepseek_v4_mtp_to_bf16.py @@ -1,16 +1,16 @@ #!/usr/bin/env python3 -"""Convert DeepSeek-V4 ``mtp.*`` checkpoint tensors to BF16. +"""Convert DeepSeek-V4 checkpoint tensors to BF16. The DeepSeek-V4-Flash checkpoint stores dense MTP matrices as block-FP8 and routed-expert matrices as packed MXFP4. Casting those tensors directly would not recover their values. This tool performs the corresponding dequantization, -removes the paired scale tensors from the output index, and writes every MTP -tensor as BF16. +removes the paired scale tensors from the output index, writes quantized and +low-precision weight tensors as BF16, and preserves source FP32 state tensors. -Non-MTP shards are linked or copied into a new model directory. The source -model is never modified. This changes checkpoint storage only; the runtime -must select the no-quantization path for the converted MTP layers while keeping -the target layers on their original quantization path. +By default only ``mtp.*`` tensors are converted and non-MTP shards are linked +or copied into the new model directory. ``--all-weights`` converts the entire +checkpoint, preserves integer routing tables, and removes the checkpoint's +quantization configuration. The source model is never modified. """ from __future__ import annotations @@ -77,6 +77,7 @@ class TensorSpec: source_dtype: str source_shape: Tuple[int, ...] output_shape: Tuple[int, ...] + output_dtype: str output_nbytes: int conversion: str scale_name: str | None = None @@ -105,13 +106,13 @@ def _parse_size(value: str) -> int: units = { "B": 1, "KB": 1000, - "MB": 1000**2, - "GB": 1000**3, - "TB": 1000**4, + "MB": 1000 ** 2, + "GB": 1000 ** 3, + "TB": 1000 ** 4, "KIB": 1024, - "MIB": 1024**2, - "GIB": 1024**3, - "TIB": 1024**4, + "MIB": 1024 ** 2, + "GIB": 1024 ** 3, + "TIB": 1024 ** 4, } for suffix in sorted(units, key=len, reverse=True): if text.endswith(suffix): @@ -149,9 +150,7 @@ def dequantize_fp8_block( rows, cols = (int(dim) for dim in output_shape) expected_scale_shape = (math.ceil(rows / block_size), math.ceil(cols / block_size)) if tuple(scale.shape) != expected_scale_shape: - raise ValueError( - f"FP8 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}" - ) + raise ValueError(f"FP8 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}") output = torch.empty((rows, cols), dtype=torch.bfloat16, device="cpu") for row_start in range(0, rows, row_chunk_size): @@ -160,16 +159,12 @@ def dequantize_fp8_block( scale_row_end = math.ceil(row_end / block_size) scale_row_offset = row_start - scale_row_start * block_size - weight_chunk = _to_device(weight_slice[row_start:row_end], device).to( - torch.float32 - ) - scale_chunk = _to_device(scale[scale_row_start:scale_row_end], device).to( - torch.float32 - ) + weight_chunk = _to_device(weight_slice[row_start:row_end], device).to(torch.float32) + scale_chunk = _to_device(scale[scale_row_start:scale_row_end], device).to(torch.float32) expanded_scale = scale_chunk.repeat_interleave(block_size, dim=0) - expanded_scale = expanded_scale[ - scale_row_offset : scale_row_offset + row_end - row_start - ].repeat_interleave(block_size, dim=1)[:, :cols] + expanded_scale = expanded_scale[scale_row_offset : scale_row_offset + row_end - row_start].repeat_interleave( + block_size, dim=1 + )[:, :cols] converted = (weight_chunk * expanded_scale).to(torch.bfloat16).cpu() output[row_start:row_end].copy_(converted) del weight_chunk, scale_chunk, expanded_scale, converted @@ -190,28 +185,20 @@ def dequantize_mxfp4( rows, packed_cols = (int(dim) for dim in packed_shape) logical_cols = packed_cols * 2 if logical_cols % block_size != 0: - raise ValueError( - f"MXFP4 logical K dimension {logical_cols} is not divisible by block size {block_size}" - ) + raise ValueError(f"MXFP4 logical K dimension {logical_cols} is not divisible by block size {block_size}") expected_scale_shape = (rows, logical_cols // block_size) if tuple(scale.shape) != expected_scale_shape: - raise ValueError( - f"MXFP4 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}" - ) + raise ValueError(f"MXFP4 scale shape mismatch: expected {expected_scale_shape}, got {tuple(scale.shape)}") output = torch.empty((rows, logical_cols), dtype=torch.bfloat16, device="cpu") lookup = torch.tensor(FP4_VALUES, dtype=torch.float32, device=device) for row_start in range(0, rows, row_chunk_size): row_end = min(row_start + row_chunk_size, rows) - packed = _to_device(packed_weight_slice[row_start:row_end], device).view( - torch.uint8 - ) + packed = _to_device(packed_weight_slice[row_start:row_end], device).view(torch.uint8) low = (packed & 0x0F).to(torch.long) high = (packed >> 4).to(torch.long) - values = torch.empty( - (row_end - row_start, logical_cols), dtype=torch.float32, device=device - ) + values = torch.empty((row_end - row_start, logical_cols), dtype=torch.float32, device=device) values[:, 0::2] = lookup[low] values[:, 1::2] = lookup[high] scale_chunk = _to_device(scale[row_start:row_end], device).to(torch.float32) @@ -237,22 +224,22 @@ def _inspect_specs( model_dir: Path, weight_map: Mapping[str, str], *, - prefix: str, + prefix: str | None, fp8_block_size: int, mxfp4_block_size: int, ) -> Tuple[List[TensorSpec], set[str], int]: - mtp_weight_map = { - name: shard for name, shard in weight_map.items() if name.startswith(prefix) + selected_weight_map = { + name: shard for name, shard in weight_map.items() if prefix is None or name.startswith(prefix) } - if not mtp_weight_map: + if not selected_weight_map: raise ValueError(f"no checkpoint tensors start with {prefix!r}") names_by_shard: Dict[str, List[str]] = {} - for name, shard in mtp_weight_map.items(): + for name, shard in selected_weight_map.items(): names_by_shard.setdefault(shard, []).append(name) tensor_meta: Dict[str, Tuple[str, Tuple[int, ...]]] = {} - original_mtp_nbytes = 0 + original_nbytes = 0 for shard, names in names_by_shard.items(): shard_path = model_dir / shard if not shard_path.is_file(): @@ -261,19 +248,17 @@ def _inspect_specs( shard_keys = set(file.keys()) missing = set(names) - shard_keys if missing: - raise ValueError( - f"{shard} is missing indexed tensors: {sorted(missing)[:5]}" - ) + raise ValueError(f"{shard} is missing indexed tensors: {sorted(missing)[:5]}") for name in names: tensor_slice = file.get_slice(name) dtype = tensor_slice.get_dtype() shape = tuple(int(dim) for dim in tensor_slice.get_shape()) tensor_meta[name] = (dtype, shape) - original_mtp_nbytes += _tensor_nbytes(dtype, shape) + original_nbytes += _tensor_nbytes(dtype, shape) paired_scales: set[str] = set() specs: List[TensorSpec] = [] - for name in mtp_weight_map: + for name in selected_weight_map: dtype, shape = tensor_meta[name] if dtype == "F8_E8M0": continue @@ -282,9 +267,7 @@ def _inspect_specs( raise ValueError(f"FP8 tensor must be 2-D: {name} has shape {shape}") scale_name = _scale_name(name) if tensor_meta.get(scale_name, (None,))[0] != "F8_E8M0": - raise ValueError( - f"missing E8M0 scale for {name}: expected {scale_name}" - ) + raise ValueError(f"missing E8M0 scale for {name}: expected {scale_name}") scale_shape = tensor_meta[scale_name][1] expected_scale_shape = ( math.ceil(shape[0] / fp8_block_size), @@ -296,17 +279,14 @@ def _inspect_specs( ) paired_scales.add(scale_name) output_shape = shape + output_dtype = "BF16" conversion = "fp8_block" elif dtype == "I8": if len(shape) != 2: - raise ValueError( - f"packed MXFP4 tensor must be 2-D: {name} has shape {shape}" - ) + raise ValueError(f"packed MXFP4 tensor must be 2-D: {name} has shape {shape}") scale_name = _scale_name(name) if tensor_meta.get(scale_name, (None,))[0] != "F8_E8M0": - raise ValueError( - f"missing E8M0 scale for {name}: expected {scale_name}" - ) + raise ValueError(f"missing E8M0 scale for {name}: expected {scale_name}") output_shape = (shape[0], shape[1] * 2) expected_scale_shape = ( output_shape[0], @@ -318,40 +298,42 @@ def _inspect_specs( f"got {tensor_meta[scale_name][1]}" ) paired_scales.add(scale_name) + output_dtype = "BF16" conversion = "mxfp4" - elif dtype in {"BF16", "F16", "F32"}: + elif dtype in {"BF16", "F16"}: scale_name = None output_shape = shape + output_dtype = "BF16" conversion = "cast" + elif dtype in {"F32", "I64"}: + scale_name = None + output_shape = shape + output_dtype = dtype + conversion = "preserve" else: - raise ValueError( - f"cannot convert {name} with dtype {dtype} to BF16 without model-specific semantics" - ) + raise ValueError(f"cannot convert {name} with dtype {dtype} to BF16 without model-specific semantics") specs.append( TensorSpec( name=name, - source_shard=mtp_weight_map[name], + source_shard=selected_weight_map[name], source_dtype=dtype, source_shape=shape, output_shape=output_shape, - output_nbytes=_numel(output_shape) * 2, + output_dtype=output_dtype, + output_nbytes=_tensor_nbytes(output_dtype, output_shape), conversion=conversion, scale_name=scale_name, ) ) - orphan_scales = { - name for name, (dtype, _) in tensor_meta.items() if dtype == "F8_E8M0" - } - paired_scales + orphan_scales = {name for name, (dtype, _) in tensor_meta.items() if dtype == "F8_E8M0"} - paired_scales if orphan_scales: - raise ValueError(f"orphan MTP E8M0 scales: {sorted(orphan_scales)[:10]}") - return specs, paired_scales, original_mtp_nbytes + raise ValueError(f"orphan E8M0 scales: {sorted(orphan_scales)[:10]}") + return specs, paired_scales, original_nbytes -def _plan_shards( - specs: Sequence[TensorSpec], max_shard_size: int -) -> List[List[TensorSpec]]: +def _plan_shards(specs: Sequence[TensorSpec], max_shard_size: int) -> List[List[TensorSpec]]: shards: List[List[TensorSpec]] = [] current: List[TensorSpec] = [] current_size = 0 @@ -415,6 +397,8 @@ def _convert_spec( ) -> torch.Tensor: if spec.conversion == "cast": return source_file.get_tensor(spec.name).to(torch.bfloat16).contiguous() + if spec.conversion == "preserve": + return source_file.get_tensor(spec.name).contiguous() assert spec.scale_name is not None weight_slice = source_file.get_slice(spec.name) @@ -440,6 +424,19 @@ def _convert_spec( raise AssertionError(f"unknown conversion: {spec.conversion}") +def _output_shard_name(prefix: str, shard_index: int, shard_count: int) -> str: + return f"{prefix}-{shard_index:05d}-of-{shard_count:05d}.safetensors" + + +def _planned_weight_map(shard_plan: Sequence[Sequence[TensorSpec]], output_shard_prefix: str) -> Dict[str, str]: + shard_count = len(shard_plan) + return { + spec.name: _output_shard_name(output_shard_prefix, shard_index, shard_count) + for shard_index, shard_specs in enumerate(shard_plan, start=1) + for spec in shard_specs + } + + def _write_converted_shards( source_dir: Path, output_dir: Path, @@ -449,21 +446,42 @@ def _write_converted_shards( mxfp4_block_size: int, row_chunk_size: int, device: torch.device, + output_shard_prefix: str, + worker_index: int, + worker_count: int, + resume: bool, ) -> Dict[str, str]: - source_shards = sorted( - {spec.source_shard for shard in shard_plan for spec in shard} - ) + assigned_shards = [ + (shard_index, shard_specs) + for shard_index, shard_specs in enumerate(shard_plan, start=1) + if (shard_index - 1) % worker_count == worker_index + ] + source_shards = sorted({spec.source_shard for _, shard_specs in assigned_shards for spec in shard_specs}) output_weight_map: Dict[str, str] = {} with ExitStack() as stack: source_files = { - shard: stack.enter_context( - safe_open(source_dir / shard, framework="pt", device="cpu") - ) + shard: stack.enter_context(safe_open(source_dir / shard, framework="pt", device="cpu")) for shard in source_shards } shard_count = len(shard_plan) - for shard_index, shard_specs in enumerate(shard_plan, start=1): - output_name = f"mtp-bf16-{shard_index:05d}-of-{shard_count:05d}.safetensors" + print( + f"worker {worker_index}/{worker_count}: assigned {len(assigned_shards)} " f"of {shard_count} shards", + flush=True, + ) + for shard_index, shard_specs in assigned_shards: + output_name = _output_shard_name(output_shard_prefix, shard_index, shard_count) + output_path = output_dir / output_name + if output_path.exists(): + if not resume: + raise FileExistsError(f"output shard already exists: {output_path}") + existing_weight_map = {spec.name: output_name for spec in shard_specs} + _verify_output(output_dir, shard_specs, existing_weight_map) + output_weight_map.update(existing_weight_map) + print( + f"[{shard_index}/{shard_count}] verified existing {output_name}", + flush=True, + ) + continue planned_bytes = sum(spec.output_nbytes for spec in shard_specs) print( f"[{shard_index}/{shard_count}] converting {len(shard_specs)} tensors " @@ -481,8 +499,7 @@ def _write_converted_shards( ) for spec in shard_specs } - temporary_path = output_dir / f".{output_name}.tmp" - output_path = output_dir / output_name + temporary_path = output_dir / f".{output_name}.{os.getpid()}.tmp" save_file(tensors, temporary_path, metadata={"format": "pt"}) os.replace(temporary_path, output_path) for spec in shard_specs: @@ -501,9 +518,7 @@ def _verify_output( ) -> None: expected_by_shard: Dict[str, Dict[str, TensorSpec]] = {} for spec in specs: - expected_by_shard.setdefault(converted_weight_map[spec.name], {})[ - spec.name - ] = spec + expected_by_shard.setdefault(converted_weight_map[spec.name], {})[spec.name] = spec for shard, expected in expected_by_shard.items(): with safe_open(output_dir / shard, framework="pt", device="cpu") as file: @@ -515,14 +530,10 @@ def _verify_output( ) for name, spec in expected.items(): tensor_slice = file.get_slice(name) - if tensor_slice.get_dtype() != "BF16": - raise ValueError( - f"{name} is {tensor_slice.get_dtype()}, expected BF16" - ) + if tensor_slice.get_dtype() != spec.output_dtype: + raise ValueError(f"{name} is {tensor_slice.get_dtype()}, expected {spec.output_dtype}") if tuple(tensor_slice.get_shape()) != spec.output_shape: - raise ValueError( - f"{name} shape is {tuple(tensor_slice.get_shape())}, expected {spec.output_shape}" - ) + raise ValueError(f"{name} shape is {tuple(tensor_slice.get_shape())}, expected {spec.output_shape}") def convert(args: argparse.Namespace) -> None: @@ -532,78 +543,95 @@ def convert(args: argparse.Namespace) -> None: raise ValueError("source and output directories must be different") if args.row_chunk_size <= 0: raise ValueError("--row-chunk-size must be positive") + if args.worker_count <= 0: + raise ValueError("--worker-count must be positive") + if not 0 <= args.worker_index < args.worker_count: + raise ValueError("--worker-index must be in [0, --worker-count)") + if args.worker_count > 1 and not args.resume: + raise ValueError("multi-worker conversion requires --resume") device = torch.device(args.device) if device.type == "cuda" and not torch.cuda.is_available(): - raise RuntimeError( - "CUDA conversion requested but torch.cuda.is_available() is false" - ) + raise RuntimeError("CUDA conversion requested but torch.cuda.is_available() is false") _, index = _load_index(source_dir) weight_map: Dict[str, str] = index["weight_map"] - specs, paired_scales, original_mtp_nbytes = _inspect_specs( + selected_prefix = None if args.all_weights else args.prefix + specs, paired_scales, original_nbytes = _inspect_specs( source_dir, weight_map, - prefix=args.prefix, + prefix=selected_prefix, fp8_block_size=args.fp8_block_size, mxfp4_block_size=args.mxfp4_block_size, ) shard_plan = _plan_shards(specs, args.max_shard_size) - output_mtp_nbytes = sum(spec.output_nbytes for spec in specs) + output_nbytes = sum(spec.output_nbytes for spec in specs) + output_shard_prefix = "model" if args.all_weights else "mtp-bf16" + converted_weight_map = _planned_weight_map(shard_plan, output_shard_prefix) conversion_counts: Dict[str, int] = {} for spec in specs: - conversion_counts[spec.conversion] = ( - conversion_counts.get(spec.conversion, 0) + 1 - ) + conversion_counts[spec.conversion] = conversion_counts.get(spec.conversion, 0) + 1 print(f"source: {source_dir}") print(f"output: {output_dir}") - print(f"MTP output tensors: {len(specs)}; removed scales: {len(paired_scales)}") + scope = "all checkpoint tensors" if args.all_weights else f"prefix {args.prefix!r}" + print(f"scope: {scope}") + print(f"output tensors: {len(specs)}; removed scales: {len(paired_scales)}") print(f"conversions: {conversion_counts}") print( - f"MTP size: {original_mtp_nbytes / 1024**3:.2f} GiB -> " - f"{output_mtp_nbytes / 1024**3:.2f} GiB in {len(shard_plan)} shards" + f"selected size: {original_nbytes / 1024**3:.2f} GiB -> " + f"{output_nbytes / 1024**3:.2f} GiB in {len(shard_plan)} shards" ) if args.dry_run: return - non_mtp_weight_map = { - name: shard - for name, shard in weight_map.items() - if not name.startswith(args.prefix) - } - _prepare_output_dir( - source_dir, - output_dir, - referenced_non_mtp_shards=non_mtp_weight_map.values(), - link_mode=args.link_mode, - ) - converted_weight_map = _write_converted_shards( - source_dir, - output_dir, - shard_plan, - fp8_block_size=args.fp8_block_size, - mxfp4_block_size=args.mxfp4_block_size, - row_chunk_size=args.row_chunk_size, - device=device, - ) + if args.resume or args.finalize_only: + if not output_dir.is_dir(): + raise FileNotFoundError(f"output directory does not exist: {output_dir}") + else: + unchanged_weight_map = { + name: shard + for name, shard in weight_map.items() + if selected_prefix is not None and not name.startswith(selected_prefix) + } + _prepare_output_dir( + source_dir, + output_dir, + referenced_non_mtp_shards=unchanged_weight_map.values(), + link_mode=args.link_mode, + ) + + if not args.finalize_only: + _write_converted_shards( + source_dir, + output_dir, + shard_plan, + fp8_block_size=args.fp8_block_size, + mxfp4_block_size=args.mxfp4_block_size, + row_chunk_size=args.row_chunk_size, + device=device, + output_shard_prefix=output_shard_prefix, + worker_index=args.worker_index, + worker_count=args.worker_count, + resume=args.resume, + ) + if args.worker_count > 1: + print(f"worker {args.worker_index}/{args.worker_count} complete") + return new_weight_map: Dict[str, str] = {} for name, shard in weight_map.items(): - if not name.startswith(args.prefix): + selected = selected_prefix is None or name.startswith(selected_prefix) + if not selected: new_weight_map[name] = shard elif name in converted_weight_map: new_weight_map[name] = converted_weight_map[name] elif name not in paired_scales: - raise AssertionError( - f"MTP tensor was neither converted nor removed: {name}" - ) + raise AssertionError(f"selected tensor was neither converted nor removed: {name}") metadata = dict(index.get("metadata") or {}) if "total_size" in metadata: - metadata["total_size"] = ( - int(metadata["total_size"]) - original_mtp_nbytes + output_mtp_nbytes - ) + metadata["total_size"] = int(metadata["total_size"]) - original_nbytes + output_nbytes new_index = {"metadata": metadata, "weight_map": new_weight_map} index_path = output_dir / "model.safetensors.index.json" temporary_index_path = output_dir / ".model.safetensors.index.json.tmp" @@ -612,22 +640,53 @@ def convert(args: argparse.Namespace) -> None: file.write("\n") os.replace(temporary_index_path, index_path) - manifest = { - "source_model_dir": str(source_dir), - "weight_prefix": args.prefix, - "output_dtype": "bfloat16", - "runtime_note": ( - "Configure converted MTP layers with quant type 'none'; target layers " - "remain on their original quantization path." - ), - "fp8_block_size": args.fp8_block_size, - "mxfp4_block_size": args.mxfp4_block_size, - "converted_tensor_count": len(specs), - "removed_scale_count": len(paired_scales), - "original_mtp_bytes": original_mtp_nbytes, - "output_mtp_bytes": output_mtp_nbytes, - } - with (output_dir / "mtp_bf16_conversion.json").open("w", encoding="utf-8") as file: + if args.all_weights: + source_config_path = source_dir / "config.json" + with source_config_path.open("r", encoding="utf-8") as file: + config = json.load(file) + config["torch_dtype"] = "bfloat16" + config.pop("quantization_config", None) + config.pop("expert_dtype", None) + temporary_config_path = output_dir / ".config.json.tmp" + with temporary_config_path.open("w", encoding="utf-8") as file: + json.dump(config, file, ensure_ascii=False, indent=2) + file.write("\n") + os.replace(temporary_config_path, output_dir / "config.json") + + if args.all_weights: + manifest = { + "source_model_dir": str(source_dir), + "conversion_scope": "all_weights", + "output_dtype": "bfloat16_with_fp32_preserved", + "source_fp32_tensors_preserved": sum(spec.source_dtype == "F32" for spec in specs), + "source_fp32_bytes_preserved": sum(spec.output_nbytes for spec in specs if spec.source_dtype == "F32"), + "integer_aux_tensors_preserved": sum(spec.source_dtype == "I64" for spec in specs), + "fp8_block_size": args.fp8_block_size, + "mxfp4_block_size": args.mxfp4_block_size, + "output_tensor_count": len(specs), + "removed_scale_count": len(paired_scales), + "original_selected_bytes": original_nbytes, + "output_selected_bytes": output_nbytes, + } + manifest_name = "bf16_conversion.json" + else: + manifest = { + "source_model_dir": str(source_dir), + "weight_prefix": args.prefix, + "output_dtype": "bfloat16_with_fp32_preserved", + "runtime_note": ( + "Configure converted MTP layers with quant type 'none'; target layers " + "remain on their original quantization path." + ), + "fp8_block_size": args.fp8_block_size, + "mxfp4_block_size": args.mxfp4_block_size, + "converted_tensor_count": len(specs), + "removed_scale_count": len(paired_scales), + "original_mtp_bytes": original_nbytes, + "output_mtp_bytes": output_nbytes, + } + manifest_name = "mtp_bf16_conversion.json" + with (output_dir / manifest_name).open("w", encoding="utf-8") as file: json.dump(manifest, file, ensure_ascii=False, indent=2) file.write("\n") @@ -641,8 +700,12 @@ def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("source_model_dir", help="source Hugging Face model directory") parser.add_argument("output_model_dir", help="new output model directory") - parser.add_argument( - "--prefix", default="mtp.", help="checkpoint key prefix to convert" + selection = parser.add_mutually_exclusive_group() + selection.add_argument("--prefix", default="mtp.", help="checkpoint key prefix to convert") + selection.add_argument( + "--all-weights", + action="store_true", + help="convert quantized/low-precision weights and preserve FP32 state and integer routing tables", ) parser.add_argument( "--device", @@ -679,6 +742,28 @@ def build_parser() -> argparse.ArgumentParser: action="store_true", help="skip the final dtype/key/shape verification", ) + parser.add_argument( + "--worker-count", + type=int, + default=1, + help="number of independent shard workers sharing the output directory", + ) + parser.add_argument( + "--worker-index", + type=int, + default=0, + help="zero-based index of this shard worker", + ) + parser.add_argument( + "--resume", + action="store_true", + help="reuse verified output shards in an existing output directory", + ) + parser.add_argument( + "--finalize-only", + action="store_true", + help="write metadata and verify all planned shards without converting", + ) return parser From 0ed251199d7df166be65974e9846599409de94f8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 28 Aug 2026 05:34:42 +0000 Subject: [PATCH 151/214] feat(disk-cache): isolate cache directories per server instance --- lightllm/server/api_cli.py | 3 ++- lightllm/server/multi_level_kv_cache/disk_cache_worker.py | 5 +++-- 2 files changed, 5 insertions(+), 3 deletions(-) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index bdb095184a..8b2785aee3 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -857,7 +857,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--disk_cache_dir", type=str, default=None, - help="""Directory used to persist disk cache data. Defaults to a temp directory when not set.""", + help="""Base directory used to persist disk cache data. A unique service name is appended so every server + instance uses a separate subdirectory. Defaults to a temp directory when not set.""", ) parser.add_argument( "--enable_dp_prompt_cache_fetch", diff --git a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py index 542ddbd877..07a0be319b 100644 --- a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py +++ b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py @@ -50,8 +50,9 @@ def __init__( # 读写同时进行时,分配8线程用来写,16线程用来读 max_concurrent_write_tasks = 8 - cache_dir = disk_cache_dir - if not cache_dir: + if disk_cache_dir: + cache_dir = os.path.join(disk_cache_dir, f"lightllm_disk_cache_{get_unique_server_name()}") + else: cache_dir = os.path.join(tempfile.gettempdir(), f"lightllm_disk_cache_{get_unique_server_name()}") os.makedirs(cache_dir, exist_ok=True) cache_file = os.path.join(cache_dir, "cache_file") From b8073f843c71dcd3f837b3b17522383dc3688b9d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 1 Sep 2026 02:17:06 +0000 Subject: [PATCH 152/214] use underlying HF tokenizer for xgramma --- lightllm/models/deepseek_v4/model.py | 4 ++++ lightllm/server/core/objs/sampling_params.py | 9 +++++++-- .../chunked_prefill/impl_for_xgrammar_mode.py | 3 ++- 3 files changed, 13 insertions(+), 3 deletions(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 5f236d2e4b..0388f86219 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -363,6 +363,10 @@ def __init__(self, tokenizer, model_dir): def __getattr__(self, name): return getattr(self.tokenizer, name) + @property + def xgrammar_tokenizer(self): + return self.tokenizer + def get_added_vocab(self): if self._added_vocab is None: self._added_vocab = self.tokenizer.get_added_vocab() diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 02987446fa..cebadc0ba5 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -23,6 +23,11 @@ MAX_SEED = (1 << 63) - 1 +def get_xgrammar_tokenizer(tokenizer): + """Return a tokenizer's explicit xgrammar-compatible tokenizer, if any.""" + return getattr(tokenizer, "xgrammar_tokenizer", tokenizer) + + class StopSequence(ctypes.Structure): _pack_ = 4 _fields_ = [ @@ -149,7 +154,7 @@ def initialize(self, constraint: str, tokenizer): if self.length > 0 and tokenizer is not None and constraint != "json": import xgrammar as xgr - tokenizer_info = xgr.TokenizerInfo.from_huggingface(tokenizer) + tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(tokenizer)) xgrammar_compiler = xgr.GrammarCompiler(tokenizer_info, max_threads=8) xgrammar_compiler.compile_grammar(constraint) except Exception as e: @@ -179,7 +184,7 @@ def initialize(self, constraint: str, tokenizer): if self.length > 0 and tokenizer is not None: import xgrammar as xgr - tokenizer_info = xgr.TokenizerInfo.from_huggingface(tokenizer) + tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(tokenizer)) xgrammar_compiler = xgr.GrammarCompiler(tokenizer_info, max_threads=8) xgrammar_compiler.compile_json_schema(constraint) except Exception as e: diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py index b159b98f25..683d345d4d 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py @@ -4,6 +4,7 @@ from .impl import ChunkedPrefillBackend from lightllm.utils.infer_utils import calculate_time from lightllm.server.core.objs import FinishStatus +from lightllm.server.core.objs.sampling_params import get_xgrammar_tokenizer from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.server.tokenizer import get_tokenizer from lightllm.utils.log_utils import init_logger @@ -26,7 +27,7 @@ def init_custom(self): self.args.model_dir, self.args.tokenizer_mode, trust_remote_code=self.args.trust_remote_code ) - self.tokenizer_info = xgr.TokenizerInfo.from_huggingface(self.tokenizer) + self.tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(self.tokenizer)) self.xgrammar_compiler = xgr.GrammarCompiler(self.tokenizer_info, max_threads=8) self.xgrammar_token_bitmask = xgr.allocate_token_bitmask(1, self.tokenizer_info.vocab_size) From e3655357b619c7ad87d679870dd2d7fae7feb79e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 2 Sep 2026 05:33:59 +0000 Subject: [PATCH 153/214] support DeepSeek-V4-Flash-Vision --- .../models/deepseek_v4/deepseek_v4_visual.py | 226 ++++++++ .../models/deepseek_v4/image_processor.py | 165 ++++++ lightllm/models/deepseek_v4/infer_struct.py | 33 +- .../layer_infer/pre_layer_infer.py | 9 +- .../layer_infer/transformer_layer_infer.py | 18 +- .../layer_weights/transformer_layer_weight.py | 5 + lightllm/models/deepseek_v4/model.py | 67 ++- .../triton_kernel/build_swa_index_dsv4.py | 86 ++- .../triton_kernel/csrc/topk_softplus_sqrt.cu | 520 ++++++++++++++++++ .../triton_kernel/topk_softplus_sqrt.py | 65 +++ lightllm/models/deepseek_v4/workspace.py | 5 +- lightllm/models/gemma4/tokenizer.py | 2 + lightllm/server/httpserver/manager.py | 8 + lightllm/server/multimodal_params.py | 6 + .../server/router/model_infer/infer_batch.py | 36 +- .../mode_backend/dsv4_multi_level_kv_cache.py | 12 + lightllm/server/tokenizer.py | 2 +- .../visualserver/model_infer/model_rpc.py | 3 + lightllm/utils/config_utils.py | 2 + unit_tests/models/deepseek_v4/test_vision.py | 294 ++++++++++ .../deepseek_v4/test_vision_integration.py | 408 ++++++++++++++ unit_tests/models/gemma4/test_tokenizer.py | 38 ++ 22 files changed, 1984 insertions(+), 26 deletions(-) create mode 100644 lightllm/models/deepseek_v4/deepseek_v4_visual.py create mode 100644 lightllm/models/deepseek_v4/image_processor.py create mode 100644 lightllm/models/deepseek_v4/triton_kernel/csrc/topk_softplus_sqrt.cu create mode 100644 lightllm/models/deepseek_v4/triton_kernel/topk_softplus_sqrt.py create mode 100644 unit_tests/models/deepseek_v4/test_vision.py create mode 100644 unit_tests/models/deepseek_v4/test_vision_integration.py create mode 100644 unit_tests/models/gemma4/test_tokenizer.py diff --git a/lightllm/models/deepseek_v4/deepseek_v4_visual.py b/lightllm/models/deepseek_v4/deepseek_v4_visual.py new file mode 100644 index 0000000000..9f0193d280 --- /dev/null +++ b/lightllm/models/deepseek_v4/deepseek_v4_visual.py @@ -0,0 +1,226 @@ +# Copyright (c) 2023 DeepSeek + +import json +import os +from functools import lru_cache +from io import BytesIO +from types import SimpleNamespace +from typing import List + +import torch +import torch.nn.functional as F +from PIL import Image +from safetensors import safe_open +from torch import nn + +from lightllm.models.deepseek_v4.image_processor import ( + IMAGE, + build_canonical_image_block, + load_image, +) +from lightllm.server.embed_cache.utils import get_shm_name_data, read_shm +from lightllm.server.multimodal_params import ImageItem + + +@lru_cache(8) +def get_vision_cos_sin(n_h: int, n_w: int, dim: int, theta: float): + inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) + hpos = torch.arange(n_h).unsqueeze(1).expand(n_h, n_w) + wpos = torch.arange(n_w).unsqueeze(0).expand(n_h, n_w) + freqs = torch.stack([hpos, wpos], dim=-1).reshape(-1, 2, 1).float() * inv_freq + freqs = freqs.flatten(1) + return freqs.cos().unsqueeze(1), freqs.sin().unsqueeze(1) + + +def apply_rotary(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x1, x2 = x.float().chunk(2, dim=-1) + return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1).to(dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim: int, eps: float = 1e-6): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + dtype = x.dtype + x = x.float() + x = x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps) + return (self.weight * x).to(dtype) + + +class PatchEmbed(nn.Module): + def __init__(self, args): + super().__init__() + self.proj = nn.Linear( + 3 * args.vision_patch_size ** 2, + args.vision_dim, + dtype=torch.bfloat16, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.proj(x.flatten(1)) + + +class Attention(nn.Module): + def __init__(self, args): + super().__init__() + self.n_heads = args.vision_n_heads + self.head_dim = args.vision_dim // args.vision_n_heads + self.wqkv = nn.Linear(args.vision_dim, 3 * args.vision_dim, dtype=torch.bfloat16) + self.wo = nn.Linear(args.vision_dim, args.vision_dim, dtype=torch.bfloat16) + + def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + n = x.size(0) + q, k, v = (t.view(n, self.n_heads, self.head_dim) for t in self.wqkv(x).chunk(3, dim=-1)) + q = apply_rotary(q, cos, sin) + k = apply_rotary(k, cos, sin) + o = F.scaled_dot_product_attention(q.transpose(0, 1), k.transpose(0, 1), v.transpose(0, 1)) + return self.wo(o.transpose(0, 1).reshape(n, -1)) + + +class MLP(nn.Module): + def __init__(self, args): + super().__init__() + self.w1 = nn.Linear( + args.vision_dim, + 2 * args.vision_inter_dim, + bias=False, + dtype=torch.bfloat16, + ) + self.w2 = nn.Linear( + args.vision_inter_dim, + args.vision_dim, + bias=False, + dtype=torch.bfloat16, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + gate, up = self.w1(x).chunk(2, dim=-1) + return self.w2(F.silu(gate) * up) + + +class Block(nn.Module): + def __init__(self, args): + super().__init__() + self.norm1 = RMSNorm(args.vision_dim) + self.attn = Attention(args) + self.norm2 = RMSNorm(args.vision_dim) + self.mlp = MLP(args) + + def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + x = x + self.attn(self.norm1(x), cos, sin) + return x + self.mlp(self.norm2(x)) + + +class ViT(nn.Module): + """DeepSeek ViT: full bidirectional attention over one image with 2D RoPE.""" + + def __init__(self, args): + super().__init__() + self.rope_dim = args.vision_dim // args.vision_n_heads // 2 + self.rope_theta = args.vision_rope_theta + self.patch_embed = PatchEmbed(args) + self.blocks = nn.ModuleList([Block(args) for _ in range(args.vision_n_layers)]) + self.norm = RMSNorm(args.vision_dim) + + def forward(self, patches: torch.Tensor, n_h: int, n_w: int) -> torch.Tensor: + x = self.patch_embed(patches) + cos, sin = get_vision_cos_sin(n_h, n_w, self.rope_dim, self.rope_theta) + cos = cos.to(device=x.device) + sin = sin.to(device=x.device) + for block in self.blocks: + x = block(x, cos, sin) + return self.norm(x) + + +class Aligner(nn.Module): + def __init__(self, args): + super().__init__() + self.downsample_ratio = args.vision_downsample_ratio + in_dim = args.vision_dim * self.downsample_ratio ** 2 + self.w1 = nn.Linear(in_dim, args.hidden_size, dtype=torch.bfloat16) + self.w2 = nn.Linear(args.hidden_size, args.hidden_size, dtype=torch.bfloat16) + + def forward(self, x: torch.Tensor, n_h: int, n_w: int) -> torch.Tensor: + r = self.downsample_ratio + x = x.view(n_h, n_w, -1).permute(2, 0, 1) + x = F.pad(x, (0, -n_w % r, 0, -n_h % r)) + x = F.unfold(x.unsqueeze(0), r, stride=r).squeeze(0).transpose(0, 1) + return self.w2(F.gelu(self.w1(x))) + + +class DeepseekV4VisionModel(nn.Module): + """Replicated visual tower and cacheable sentinel-block encoder.""" + + _WEIGHT_PREFIXES = ("vision.", "aligner.", "image_") + + def __init__(self, kvargs, **config): + super().__init__() + self.config = SimpleNamespace(**config) + self.vision = ViT(self.config) + self.aligner = Aligner(self.config) + self.image_start = nn.Parameter(torch.empty(self.config.hidden_size, dtype=torch.bfloat16)) + self.image_end = nn.Parameter(torch.empty(self.config.hidden_size, dtype=torch.bfloat16)) + self.image_newline = nn.Parameter(torch.empty(self.config.hidden_size, dtype=torch.bfloat16)) + self.image_pad = nn.Parameter(torch.empty(self.config.hidden_size, dtype=torch.bfloat16)) + + def load_model(self, weight_dir): + index_path = os.path.join(weight_dir, "model.safetensors.index.json") + with open(index_path, "r") as index_file: + weight_map = json.load(index_file)["weight_map"] + + keys_by_file = {} + for key, file_name in weight_map.items(): + if key.startswith(self._WEIGHT_PREFIXES): + keys_by_file.setdefault(file_name, []).append(key) + + state_dict = {} + for file_name, keys in sorted(keys_by_file.items()): + with safe_open(os.path.join(weight_dir, file_name), framework="pt", device="cpu") as weight_file: + for key in keys: + state_dict[key] = weight_file.get_tensor(key) + self.load_state_dict(state_dict) + return self + + @torch.inference_mode() + def encode(self, images: List[ImageItem]): + blocks = [] + uuids = [] + valid_ids = [] + valid_start = 0 + sentinel_table = torch.stack( + [ + self.image_start, + self.image_pad, + self.image_pad, + self.image_newline, + self.image_end, + ] + ) + + for image_item in images: + uuids.append(image_item.uuid) + image_data = read_shm(get_shm_name_data(image_item.uuid)) + with Image.open(BytesIO(image_data)) as image: + patches, n_vit_h, n_vit_w, n_llm_h, n_llm_w = load_image(image, self.config) + + patches = patches.to(device=self.image_start.device, non_blocking=True) + row_major_embeds = self.aligner( + self.vision(patches, n_vit_h, n_vit_w), + n_vit_h, + n_vit_w, + ) + types, perm = build_canonical_image_block(n_llm_h, n_llm_w) + types = types.to(device=row_major_embeds.device) + block = sentinel_table[types] + block[types == IMAGE] = row_major_embeds[perm.to(device=row_major_embeds.device)] + blocks.append(block) + + valid_end = valid_start + block.shape[0] + valid_ids.append([valid_start, valid_end]) + valid_start = valid_end + + return torch.cat(blocks), uuids, valid_ids diff --git a/lightllm/models/deepseek_v4/image_processor.py b/lightllm/models/deepseek_v4/image_processor.py new file mode 100644 index 0000000000..ea61a10584 --- /dev/null +++ b/lightllm/models/deepseek_v4/image_processor.py @@ -0,0 +1,165 @@ +# Copyright (c) 2023 DeepSeek + +import math +from typing import Optional, Tuple + +import numpy as np +import torch +from PIL import Image, ImageOps + + +IMAGE_START, IMAGE_PAD, IMAGE, IMAGE_NEW_LINE, IMAGE_END = range(5) +COMPRESS_PAD_TO = 4 + + +def grid_tokens(best_height, best_width, patch_size, downsample_ratio): + """Number of LLM tokens occupied by an N-layout image block without compressor padding.""" + n_llm_h = math.ceil((best_height // patch_size) / downsample_ratio) + n_llm_w = math.ceil((best_width // patch_size) / downsample_ratio) + num_tokens = n_llm_h * (n_llm_w + 1) + 2 + if n_llm_h % 2 == 1: + num_tokens += n_llm_w + 1 + num_tokens += (n_llm_h + 1) // 2 * (n_llm_w + 1) % 2 * 2 + return n_llm_h, n_llm_w, num_tokens + + +def solve_resize_ratio(height, width, patch_size, downsample_ratio, max_n_token): + ratio = height / width + max_w_float = math.sqrt((max_n_token - 2) / ratio + 0.25) - 0.5 + max_h_float = max_w_float * ratio + if max_w_float < 1.0: + max_w = 1 + max_h = (max_n_token - 2) // (max_w + 1) + if max_h % 2 == 1: + max_h -= 1 + best_width = max_w * patch_size * downsample_ratio + best_height = max_h * patch_size * downsample_ratio + elif max_h_float < 2.0: + max_h = 2 + max_w = ((max_n_token - 2) // max_h) - 1 + assert max_w > 1 + best_width = max_w * patch_size * downsample_ratio + best_height = max_h * patch_size * downsample_ratio + else: + max_w = math.floor(max_w_float) + max_h = math.floor(max_h_float) + if max_h % 2 == 1: + max_h -= 1 + beta = min( + max_w * patch_size * downsample_ratio / width, + max_h * patch_size * downsample_ratio / height, + ) + best_width = math.floor(width * beta / patch_size) * patch_size + best_height = math.floor(height * beta / patch_size) * patch_size + n_llm_h, n_llm_w, num_tokens = grid_tokens(best_height, best_width, patch_size, downsample_ratio) + return n_llm_h, n_llm_w, best_height, best_width, num_tokens + + +def safe_resize(height, width, best_height, best_width, patch_size, downsample_ratio, max_n_token): + max_n_token -= COMPRESS_PAD_TO - 1 + n_llm_h, n_llm_w, num_tokens = grid_tokens(best_height, best_width, patch_size, downsample_ratio) + budget = max_n_token + while num_tokens > max_n_token: + n_llm_h, n_llm_w, best_height, best_width, num_tokens = solve_resize_ratio( + height, width, patch_size, downsample_ratio, budget + ) + budget -= 1 + return n_llm_h, n_llm_w, best_height, best_width + + +def get_image_grid( + height: int, + width: int, + patch_size: int, + downsample_ratio: int, + max_n_token: int, + min_pixels: int, + max_wh_ratio: Optional[float], +) -> Tuple[int, int, int, int, int]: + """Resolve ViT/LLM grids and the position-independent canonical block length.""" + if max_wh_ratio is not None and width > height * max_wh_ratio: + width = height * max_wh_ratio + if 0 < width * height < min_pixels: + ratio = (min_pixels / (width * height)) ** 0.5 + width = int(width * ratio) + height = int(height * ratio) + + best_width = math.ceil(width / patch_size) * patch_size + best_height = math.ceil(height / patch_size) * patch_size + n_llm_h, n_llm_w, best_height, best_width = safe_resize( + height, + width, + best_height, + best_width, + patch_size, + downsample_ratio, + max_n_token, + ) + n_vit_h = best_height // patch_size + n_vit_w = best_width // patch_size + _, _, core_token_num = grid_tokens(best_height, best_width, patch_size, downsample_ratio) + canonical_token_num = COMPRESS_PAD_TO - 1 + core_token_num + return n_vit_h, n_vit_w, n_llm_h, n_llm_w, canonical_token_num + + +def load_image(image: Image.Image, args): + """Transform one image into BF16 ViT patches using the checkpoint preprocessing contract.""" + image = image.convert("RGB") + n_vit_h, n_vit_w, n_llm_h, n_llm_w, _ = get_image_grid( + image.height, + image.width, + args.vision_patch_size, + args.vision_downsample_ratio, + args.vision_max_n_token, + args.vision_min_pixels, + args.vision_max_wh_ratio, + ) + best_height = n_vit_h * args.vision_patch_size + best_width = n_vit_w * args.vision_patch_size + if args.vision_max_wh_ratio is not None and image.width >= args.vision_max_wh_ratio * image.height: + image = image.resize((best_width, best_height)) + else: + image = ImageOps.pad(image, (best_width, best_height), color=(127, 127, 127)) + + x = torch.from_numpy(np.asarray(image, dtype=np.float32)).permute(2, 0, 1) / 255 + x = ((x - 0.5) / 0.5).to(torch.bfloat16) + patch_size = args.vision_patch_size + patches = ( + x.reshape(3, n_vit_h, patch_size, n_vit_w, patch_size) + .permute(1, 3, 0, 2, 4) + .reshape(n_vit_h * n_vit_w, 3, patch_size, patch_size) + ) + return patches, n_vit_h, n_vit_w, n_llm_h, n_llm_w + + +def build_image_block(n_llm_h: int, n_llm_w: int, start_pos: int): + """Build final N-layout token types and the row-major aligner permutation.""" + compress_pad = COMPRESS_PAD_TO - 1 - start_pos % COMPRESS_PAD_TO + pad_h = n_llm_h % 2 + rows = n_llm_h + pad_h + row_len = n_llm_w + 1 + pad_last = rows // 2 * row_len % 2 * 2 + types = torch.tensor( + ([IMAGE] * n_llm_w + [IMAGE_NEW_LINE]) * n_llm_h + [IMAGE_PAD] * (row_len * pad_h), + dtype=torch.int64, + ) + order = torch.arange(rows * row_len).view(rows // 2, 2, row_len).transpose(1, 2).reshape(-1) + image_idx = torch.full((rows * row_len,), -1, dtype=torch.int64) + image_idx.view(rows, row_len)[:n_llm_h, :n_llm_w] = torch.arange(n_llm_h * n_llm_w).view(n_llm_h, n_llm_w) + perm = image_idx[order] + perm = perm[perm >= 0] + types = torch.cat( + [ + torch.full((compress_pad,), IMAGE_PAD, dtype=torch.int64), + torch.tensor([IMAGE_START]), + types[order], + torch.full((pad_last,), IMAGE_PAD, dtype=torch.int64), + torch.tensor([IMAGE_END]), + ] + ) + return types, perm + + +def build_canonical_image_block(n_llm_h: int, n_llm_w: int): + """Build the cacheable image block with its maximum three compressor-alignment pads.""" + return build_image_block(n_llm_h, n_llm_w, start_pos=0) diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index 488cf82f92..05e573f7ce 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -23,6 +23,8 @@ def __init__(self): self.dsv4_sparse_req_idx = None self.dsv4_swa_indices = None self.dsv4_swa_lengths = None + self.dsv4_image_left = None + self.dsv4_image_right = None self.dsv4_c128_indices = None self.dsv4_c128_lengths = None self.dsv4_workspace = None @@ -63,11 +65,35 @@ def init_some_extra_state(self, model): self.dsv4_sparse_req_idx = self.b_req_idx self._dsv4_token_to_batch_idx = None # Sliding-window indices are layer-independent, so build them once into the model workspace. - from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_swa_index + from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import ( + build_image_visibility, + build_swa_index, + ) workspace = model.dsv4_workspace self.dsv4_workspace = workspace - self.dsv4_swa_indices, self.dsv4_swa_lengths = workspace.swa(self.microbatch_index, pos.numel()) + has_vision = model.has_vision + swa_width = workspace.vision_swa_width if has_vision and self.is_prefill else workspace.sliding_window + self.dsv4_swa_indices, self.dsv4_swa_lengths = workspace.swa( + self.microbatch_index, pos.numel(), width=swa_width + ) + if has_vision and self.is_prefill: + self.dsv4_image_left = torch.zeros_like(pos, dtype=torch.int32) + self.dsv4_image_right = torch.zeros_like(pos, dtype=torch.int32) + image_spans = [] + for batch_id, params in enumerate(self.multimodal_params): + for image in params["images"]: + # start_idx points at IMAGE_START; token_num also includes three canonical alignment slots. + image_spans.append((batch_id, image["start_idx"], image["token_num"] - 3)) + if image_spans: + build_image_visibility( + image_spans=torch.tensor(image_spans, dtype=torch.int32).cuda(non_blocking=True), + b_q_start_loc=self.b_q_start_loc, + b_ready_cache_len=self.b_ready_cache_len, + b_q_seq_len=self.b_q_seq_len, + image_left=self.dsv4_image_left, + image_right=self.dsv4_image_right, + ) self.dsv4_swa_indices, self.dsv4_swa_lengths = build_swa_index( req_idx=self.dsv4_sparse_req_idx, positions=self.position_ids, @@ -75,6 +101,9 @@ def init_some_extra_state(self, model): full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, swa_index=self.dsv4_swa_indices, swa_length=self.dsv4_swa_lengths, + window=workspace.sliding_window, + image_left=self.dsv4_image_left, + image_right=self.dsv4_image_right, ) from lightllm.models.deepseek_v4.triton_kernel.build_compress_index_dsv4 import build_compress_index diff --git a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py index b95f5a14a8..4f00af7fc6 100644 --- a/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py @@ -1,20 +1,25 @@ import torch import torch.distributed as dist from lightllm.models.llama.layer_infer.pre_layer_infer import LlamaPreLayerInfer +from lightllm.models.qwen_vl.layer_infer.pre_layer_infer import LlamaMultimodalPreLayerInfer from lightllm.distributed.communication_op import all_reduce from ..infer_struct import DeepseekV4InferStateInfo -class DeepseekV4PreLayerInfer(LlamaPreLayerInfer): +class DeepseekV4PreLayerInfer(LlamaMultimodalPreLayerInfer): """Token embedding, then expand to the hc_mult parallel residual streams [T, hc_mult*hidden].""" def __init__(self, network_config): super().__init__(network_config) self.hc_mult = network_config["hc_mult"] + self.has_vision = network_config.get("vision_n_layers", 0) > 0 return def context_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, layer_weight): - input_embdings = super().context_forward(input_ids, infer_state, layer_weight) + if self.has_vision: + input_embdings = super().context_forward(input_ids, infer_state, layer_weight) + else: + input_embdings = LlamaPreLayerInfer.context_forward(self, input_ids, infer_state, layer_weight) t, hidden = input_embdings.shape return input_embdings.unsqueeze(1).expand(t, self.hc_mult, hidden).reshape(t, self.hc_mult * hidden) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index ecb147e5a2..166cc0efb0 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -9,7 +9,6 @@ from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.tensor_utils import tensor_to_no_ref_tensor -from lightllm.utils.vllm_utils import vllm_ops from .hyper_connection import hc_pre, hc_fused_post_pre, hc_post from .compressor import fused_compress as fused_compress_op from .compressor import apply_ape @@ -17,6 +16,7 @@ from ..infer_struct import DeepseekV4InferStateInfo import deep_gemm from lightllm.models.deepseek_v4.triton_kernel.topk_transform import topk_transform_512 +from lightllm.models.deepseek_v4.triton_kernel.topk_softplus_sqrt import topk_softplus_sqrt _C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 @@ -38,6 +38,8 @@ def __init__(self, layer_num, network_config): self.hc_eps = network_config["hc_eps"] self.compress_ratio = network_config["compress_ratios"][layer_num] self.is_hash = layer_num < network_config["num_hash_layers"] + self.has_vision = network_config.get("vision_n_layers", 0) > 0 + self.vocab_size = network_config["vocab_size"] self.is_last_layer = layer_num == network_config["n_layer"] - 1 # complex64 rope table for this layer's variant (sliding / compressed); set by # DeepseekV4TpPartModel._init_to_get_rotary once the tables are built. The full compress @@ -567,8 +569,10 @@ def _select_experts( ): M = logits.shape[0] bias = None + bias_vl = None input_tokens = None hash_indices_table = None + image_token_start = 0 indices_dtype = torch.int64 if self.is_hash: hash_indices_table = layer_weight.gate_tid2eid_.weight @@ -576,20 +580,24 @@ def _select_experts( input_tokens = infer_state.input_ids.to(dtype=indices_dtype) else: bias = layer_weight.gate_bias_.weight + if self.has_vision and infer_state.is_prefill: + bias_vl = layer_weight.gate_bias_vl_.weight + image_token_start = self.vocab_size + if input_tokens is None: + input_tokens = infer_state.input_ids weights = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.float32, device=logits.device) indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=indices_dtype, device=logits.device) - token_expert_indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.int32, device=logits.device) - vllm_ops.topk_hash_softplus_sqrt( + topk_softplus_sqrt( weights, indices, - token_expert_indices, logits, - True, self.routed_scaling_factor, bias, input_tokens, hash_indices_table, + bias_vl, + image_token_start, ) return weights, indices diff --git a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py index 493fb1cf27..9bfb67214f 100644 --- a/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py +++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py @@ -46,6 +46,7 @@ def _parse_config(self): self.has_compressor = self.compress_ratio != 0 self.has_indexer = self.compress_ratio == 4 self.is_hash = self.layer_num_ < self.num_hash_layers + self.has_vision = cfg.get("vision_n_layers", 0) > 0 assert self.n_heads % self.tp_world_size_ == 0 assert self.o_groups % self.tp_world_size_ == 0 self.prefix = f"layers.{self.layer_num_}" @@ -213,6 +214,10 @@ def _init_moe(self): self.gate_bias_ = ParameterWeight( weight_name=f"{p}.gate.bias", data_type=torch.float32, weight_shape=(self.n_routed_experts,) ) + if self.has_vision: + self.gate_bias_vl_ = ParameterWeight( + weight_name=f"{p}.gate.bias_vl", data_type=torch.float32, weight_shape=(self.n_routed_experts,) + ) # shared expert (dense, bf16 after de-quant): w1=gate, w3=up fused (row), w2=down (col). # Named gate_up_proj/down_proj so the inherited Llama `_ffn_tp` (fused gate_up matmul + # silu_and_mul triton kernel, no swiglu clamp) drives it directly. Order [w1, w3] = [gate, up] diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 0388f86219..a4070dcf05 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -54,6 +54,7 @@ class DeepseekV4TpPartModel(LlamaTpPartModel): req_manager: DeepseekV4ReqManager mem_manager: DeepseekV4MemoryManager + has_vision = False pre_and_post_weight_class = DeepseekV4PreAndPostLayerWeight transformer_weight_class = DeepseekV4TransformerLayerWeight @@ -346,6 +347,15 @@ def build_inv_freq(base): return +@ModelRegistry( + "deepseek_v4", + is_multimodal=True, + condition=lambda cfg: cfg.get("vision_n_layers", 0) > 0, +) +class DeepseekV4VisionTpPartModel(DeepseekV4TpPartModel): + has_vision = True + + class DeepSeekV4Tokenizer: """Tokenizer wrapper for DeepSeek-V4's Python prompt encoding.""" @@ -354,9 +364,12 @@ class DeepSeekV4Tokenizer: # so advertise thinking support explicitly for tokenizer_supports_force_thinking(). supports_thinking = True - def __init__(self, tokenizer, model_dir): + def __init__(self, tokenizer, model_dir, model_config=None): self.tokenizer = tokenizer self.model_dir = model_dir + self.model_config = {} if model_config is None else model_config + self.has_vision = self.model_config.get("vision_n_layers", 0) > 0 + self.image_token_id = tokenizer.convert_tokens_to_ids("<|deepseek_image|>") if self.has_vision else None self._encoding_module = None self._added_vocab = None @@ -372,6 +385,58 @@ def get_added_vocab(self): self._added_vocab = self.tokenizer.get_added_vocab() return self._added_vocab + def init_imageitem_extral_params(self, img, multi_params, sampling_params): + return + + def get_image_token_length(self, img): + from lightllm.models.deepseek_v4.image_processor import get_image_grid + + cfg = self.model_config + _, _, _, _, token_num = get_image_grid( + img.image_h, + img.image_w, + cfg["vision_patch_size"], + cfg["vision_downsample_ratio"], + cfg["vision_max_n_token"], + cfg["vision_min_pixels"], + cfg.get("vision_max_wh_ratio"), + ) + return token_num + + def encode(self, prompt, multimodal_params=None, **kwargs): + if isinstance(prompt, str): + origin_ids = self.tokenizer.encode(prompt, **kwargs) + elif isinstance(prompt, list): + origin_ids = prompt + else: + raise ValueError(f"Unsupported prompt type: {type(prompt)}") + + images = [] if multimodal_params is None else multimodal_params.images + if not images: + return origin_ids + + placeholder_count = sum(token_id == self.image_token_id for token_id in origin_ids) + if placeholder_count != len(images): + raise ValueError(f"invalid image placeholder count: {placeholder_count} vs {len(images)}") + + input_ids = [] + image_index = 0 + for token_id in origin_ids: + if token_id != self.image_token_id: + input_ids.append(token_id) + continue + + image = images[image_index] + cache_skip = len(input_ids) % 4 + compress_pad = 3 - cache_skip + image.block_start_idx = len(input_ids) + image.start_idx = len(input_ids) + compress_pad + input_ids.extend(range(image.token_id + cache_skip, image.token_id + image.token_num)) + image.block_end_idx = len(input_ids) + image_index += 1 + + return input_ids + def _get_encoding_module(self): if self._encoding_module is not None: return self._encoding_module diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py index a1ef2d5be1..a3ca231b0e 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py @@ -3,6 +3,54 @@ import triton.language as tl +@triton.jit +def _build_image_visibility_kernel( + image_spans_ptr, # (num_images, 3): batch index, absolute image start, visible image length + b_q_start_loc_ptr, + b_ready_cache_len_ptr, + b_q_seq_len_ptr, + image_left_ptr, + image_right_ptr, + BLOCK_SIZE: tl.constexpr, +): + image_idx = tl.program_id(0) + span_ptr = image_spans_ptr + image_idx * 3 + batch_idx = tl.load(span_ptr) + image_start = tl.load(span_ptr + 1) + image_len = tl.load(span_ptr + 2) + cache_len = tl.load(b_ready_cache_len_ptr + batch_idx) + q_seq_len = tl.load(b_q_seq_len_ptr + batch_idx) + flat_image_start = tl.load(b_q_start_loc_ptr + batch_idx) + image_start - cache_len + + for start in range(0, image_len, BLOCK_SIZE): + offset = start + tl.arange(0, BLOCK_SIZE) + q_offset = image_start - cache_len + offset + mask = (offset < image_len) & (q_offset >= 0) & (q_offset < q_seq_len) + output_offset = flat_image_start + offset + tl.store(image_left_ptr + output_offset, offset, mask=mask) + tl.store(image_right_ptr + output_offset, image_len - 1 - offset, mask=mask) + + +def build_image_visibility( + image_spans: torch.Tensor, + b_q_start_loc: torch.Tensor, + b_ready_cache_len: torch.Tensor, + b_q_seq_len: torch.Tensor, + image_left: torch.Tensor, + image_right: torch.Tensor, +): + """Scatter image-relative left/right distances into the flattened prefill layout.""" + _build_image_visibility_kernel[(image_spans.shape[0],)]( + image_spans, + b_q_start_loc, + b_ready_cache_len, + b_q_seq_len, + image_left, + image_right, + BLOCK_SIZE=64, + ) + + @triton.jit def _build_swa_index_kernel( req_idx_ptr, @@ -10,10 +58,14 @@ def _build_swa_index_kernel( req_to_token_ptr, req_to_token_stride0, full_to_swa_ptr, + image_left_ptr, + image_right_ptr, swa_index_ptr, swa_index_stride0, swa_length_ptr, WINDOW: tl.constexpr, + WIDTH: tl.constexpr, + HAS_IMAGE: tl.constexpr, BLOCK_W: tl.constexpr, ): token_idx = tl.program_id(0) @@ -21,18 +73,26 @@ def _build_swa_index_kernel( pos = tl.load(pos_ptr + token_idx).to(tl.int64) w = tl.arange(0, BLOCK_W) - w_mask = w < WINDOW - # most-recent-first window, identical to the eager _swa_indices (offset = position - arange). + w_mask = w < WIDTH offset = pos - w - valid = (offset >= 0) & w_mask + valid = (offset >= 0) & (w < WINDOW) + length = tl.minimum(tl.maximum(pos + 1, 1), WINDOW) + if HAS_IMAGE: + left = tl.load(image_left_ptr + token_idx) + right = tl.load(image_right_ptr + token_idx) + is_image = (left > 0) | (right > 0) + image_start = tl.maximum(pos - (WINDOW - 1) - tl.maximum(left - (WINDOW - 1), 0), 0) + image_length = pos + right - image_start + 1 + offset = tl.where(is_image, image_start + w, offset) + valid = tl.where(is_image, w < image_length, valid) + length = tl.where(is_image, image_length, length) safe_offset = tl.where(valid, offset, 0) full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_offset, mask=valid, other=0).to(tl.int64) swa_slot = tl.load(full_to_swa_ptr + full_slot, mask=valid, other=-1) out = tl.where(valid, swa_slot, -1).to(tl.int32) tl.store(swa_index_ptr + token_idx * swa_index_stride0 + w, out, mask=w_mask) - length = tl.minimum(tl.maximum(pos + 1, 1), WINDOW).to(tl.int32) - tl.store(swa_length_ptr + token_idx, length) + tl.store(swa_length_ptr + token_idx, length.to(tl.int32)) def build_swa_index( @@ -42,6 +102,9 @@ def build_swa_index( full_to_swa_indexs: torch.Tensor, swa_index: torch.Tensor, swa_length: torch.Tensor, + window: int = None, + image_left: torch.Tensor = None, + image_right: torch.Tensor = None, ): """Per-token sliding-window FlashMLA index table, built ONCE per forward (layer-independent: full_to_swa is a single global map and the window is a model constant, so every layer's swa @@ -53,20 +116,29 @@ def build_swa_index( the reader adds the s_q axis via unsqueeze(1). """ T = positions.shape[0] - window = swa_index.shape[1] + width = swa_index.shape[1] + window = width if window is None else int(window) + assert window <= width + assert (image_left is None) == (image_right is None) if T == 0: return swa_index, swa_length + image_left_arg = positions if image_left is None else image_left + image_right_arg = positions if image_right is None else image_right _build_swa_index_kernel[(T,)]( req_idx, positions, req_to_token_indexs, req_to_token_indexs.stride(0), full_to_swa_indexs, + image_left_arg, + image_right_arg, swa_index, swa_index.stride(0), swa_length, WINDOW=window, - BLOCK_W=triton.next_power_of_2(window), + WIDTH=width, + HAS_IMAGE=image_left is not None, + BLOCK_W=triton.next_power_of_2(width), num_warps=4, ) return swa_index, swa_length diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_softplus_sqrt.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_softplus_sqrt.cu new file mode 100644 index 0000000000..a0407f492e --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_softplus_sqrt.cu @@ -0,0 +1,520 @@ +/* + * Adapted from + * https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu + * Copyright (c) 2024, The vLLM team. + * SPDX-FileCopyrightText: Copyright (c) 1993-2023 NVIDIA CORPORATION & + * AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// Adapted for LightLLM from vLLM commit +// 36d2a5086bae12d0c0b311607373e4fc2d1036aa: +// csrc/libtorch_stable/moe/topk_softplus_sqrt_kernels.cu. +// LightLLM-local CUDA binding for DeepSeek-V4's fixed top-6 router. The +// vLLM warp kernel is shared by hash and bias routing here; image tokens use +// LightLLM's multimodal-ID contract (all IDs >= image_token_start). + +#include +#include +#include +#include + +#include + +#include +#include + +namespace { + +constexpr int kTopK = 6; +constexpr int kWarpsPerBlock = 4; +constexpr int kThreadsPerBlock = 32 * kWarpsPerBlock; +constexpr int kMaxExpertsPerLane = 12; + +__device__ __forceinline__ float softplus_sqrt(float value) { + float score = sqrtf(fmaxf(value, 0.0f) + __logf(1.0f + __expf(-fabsf(value)))); + return isnan(score) ? 0.0f : score; +} + +template +__device__ __forceinline__ void pdl_wait_primary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) { + asm volatile("griddepcontrol.wait;" ::: "memory"); + } +#endif +} + +template +__device__ __forceinline__ void pdl_trigger_secondary() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (UsePDL) { + asm volatile("griddepcontrol.launch_dependents;" :::); + } +#endif +} + +template +__launch_bounds__(kThreadsPerBlock) __global__ void topk_softplus_sqrt_kernel( + const float* logits, + float* output_weights, + int64_t* output_indices, + int num_rows, + int num_experts, + float routed_scaling_factor, + const float* correction_bias, + const int64_t* input_ids, + const int64_t* tid2eid, + const float* bias_vl, + int64_t image_token_start) { + const int row = (blockIdx.x * blockDim.x + threadIdx.x) / 32; + const int lane = threadIdx.x % 32; + if (row >= num_rows) return; + + pdl_wait_primary(); + + const int64_t token_id = input_ids == nullptr ? 0 : input_ids[row]; + const bool is_image = bias_vl != nullptr && token_id >= image_token_start; + const bool use_hash = tid2eid != nullptr && !is_image; + + int expert = 0; + float weight = 0.0f; + if (use_hash) { + if (lane < kTopK) { + expert = static_cast(tid2eid[token_id * kTopK + lane]); + weight = softplus_sqrt(logits[row * num_experts + expert]); + } + } else { + const float* selection_bias = is_image ? bias_vl : correction_bias; + const int experts_per_lane = num_experts / 32; + float scores[kMaxExpertsPerLane]; + float selection_scores[kMaxExpertsPerLane]; + +#pragma unroll + for (int i = 0; i < kMaxExpertsPerLane; ++i) { + if (i < experts_per_lane) { + const int expert_id = lane + 32 * i; + const float score = softplus_sqrt(logits[row * num_experts + expert_id]); + scores[i] = score; + selection_scores[i] = score + selection_bias[expert_id]; + } else { + scores[i] = 0.0f; + selection_scores[i] = -INFINITY; + } + } + + int selected_experts[kTopK]; + float selected_weights[kTopK]; +#pragma unroll + for (int slot = 0; slot < kTopK; ++slot) { + float best_score = -INFINITY; + int best_expert = num_experts; +#pragma unroll + for (int i = 0; i < kMaxExpertsPerLane; ++i) { + const int expert_id = lane + 32 * i; + if (selection_scores[i] > best_score) { + best_score = selection_scores[i]; + best_expert = expert_id; + } + } +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + const float other_score = __shfl_xor_sync(0xffffffffu, best_score, mask); + const int other_expert = __shfl_xor_sync(0xffffffffu, best_expert, mask); + if (other_score > best_score || (other_score == best_score && other_expert < best_expert)) { + best_score = other_score; + best_expert = other_expert; + } + } + + float selected_weight = 0.0f; +#pragma unroll + for (int i = 0; i < kMaxExpertsPerLane; ++i) { + if (lane + 32 * i == best_expert) { + selected_weight = scores[i]; + selection_scores[i] = -INFINITY; + } + } +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + selected_weight += __shfl_xor_sync(0xffffffffu, selected_weight, mask); + } + selected_experts[slot] = best_expert; + selected_weights[slot] = selected_weight; + } + + if (lane < kTopK) { + expert = selected_experts[lane]; + weight = selected_weights[lane]; + } + } + + float weight_sum = weight; +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + weight_sum += __shfl_xor_sync(0xffffffffu, weight_sum, mask); + } + + pdl_trigger_secondary(); + + if (lane < kTopK) { + const int offset = row * kTopK + lane; + output_weights[offset] = weight * routed_scaling_factor / (weight_sum > 0.0f ? weight_sum : 1.0f); + output_indices[offset] = static_cast(expert); + } +} + +template +__launch_bounds__(kThreadsPerBlock) __global__ void topk_bias_softplus_sqrt_kernel( + const float* logits, + float* output_weights, + int64_t* output_indices, + int num_rows, + float routed_scaling_factor, + const float* correction_bias, + const int64_t* input_ids, + const float* bias_vl, + int64_t image_token_start) { + constexpr int kValuesPerThread = NumExperts / 32; + constexpr int kVectorsPerThread = kValuesPerThread / 4; + const int row = blockIdx.x * kWarpsPerBlock + threadIdx.y; + const int lane = threadIdx.x; + if (row >= num_rows) return; + + pdl_wait_primary(); + + const int64_t token_id = input_ids == nullptr ? 0 : input_ids[row]; + const bool is_image = bias_vl != nullptr && token_id >= image_token_start; + const float* selection_bias = is_image ? bias_vl : correction_bias; + const float4* row_logits = reinterpret_cast(logits + row * NumExperts); + float values[kValuesPerThread]; + +#pragma unroll + for (int vector = 0; vector < kVectorsPerThread; ++vector) { + const float4 packed = row_logits[vector * 32 + lane]; + const float packed_values[4] = {packed.x, packed.y, packed.z, packed.w}; +#pragma unroll + for (int i = 0; i < 4; ++i) { + const int local = vector * 4 + i; + const int expert_id = vector * 128 + lane * 4 + i; + values[local] = softplus_sqrt(packed_values[i]) + selection_bias[expert_id]; + } + } + + float selected_sum = 0.0f; +#pragma unroll + for (int slot = 0; slot < kTopK; ++slot) { + float best_score = values[0]; + int best_expert = lane * 4; +#pragma unroll + for (int local = 1; local < kValuesPerThread; ++local) { + const int vector = local / 4; + const int element = local % 4; + const int expert_id = vector * 128 + lane * 4 + element; + if (values[local] > best_score) { + best_score = values[local]; + best_expert = expert_id; + } + } +#pragma unroll + for (int mask = 16; mask > 0; mask >>= 1) { + const float other_score = __shfl_xor_sync(0xffffffffu, best_score, mask); + const int other_expert = __shfl_xor_sync(0xffffffffu, best_expert, mask); + if (other_score > best_score || (other_score == best_score && other_expert < best_expert)) { + best_score = other_score; + best_expert = other_expert; + } + } + + if (lane == 0) { + const float selected_weight = best_score - selection_bias[best_expert]; + const int offset = row * kTopK + slot; + output_weights[offset] = selected_weight; + output_indices[offset] = static_cast(best_expert); + selected_sum += selected_weight; + } + + if (slot + 1 < kTopK) { + const int vector = best_expert / 128; + const int owner_lane = (best_expert / 4) % 32; + if (lane == owner_lane) { + values[vector * 4 + best_expert % 4] = -INFINITY; + } + } + } + + if (lane == 0) { + const float scale = routed_scaling_factor / (selected_sum > 0.0f ? selected_sum : 1.0f); +#pragma unroll + for (int slot = 0; slot < kTopK; ++slot) { + output_weights[row * kTopK + slot] *= scale; + } + } + + pdl_trigger_secondary(); +} + +template +void launch_topk_bias_softplus_sqrt( + const at::Tensor& topk_weights, + const at::Tensor& topk_indices, + const at::Tensor& logits, + double routed_scaling_factor, + const float* correction_bias, + const int64_t* input_ids, + const float* bias_vl, + int64_t image_token_start, + cudaStream_t stream) { + cudaLaunchConfig_t config{}; + config.gridDim = (logits.size(0) + kWarpsPerBlock - 1) / kWarpsPerBlock; + config.blockDim = dim3(32, kWarpsPerBlock); + config.stream = stream; + + cudaLaunchAttribute attribute{}; + if constexpr (UsePDL) { + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = true; + config.attrs = &attribute; + config.numAttrs = 1; + } + + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, + topk_bias_softplus_sqrt_kernel, + logits.data_ptr(), + topk_weights.data_ptr(), + topk_indices.data_ptr(), + static_cast(logits.size(0)), + static_cast(routed_scaling_factor), + correction_bias, + input_ids, + bias_vl, + image_token_start)); +} + +template +void launch_topk_softplus_sqrt( + const at::Tensor& topk_weights, + const at::Tensor& topk_indices, + const at::Tensor& logits, + double routed_scaling_factor, + const float* correction_bias, + const int64_t* input_ids, + const int64_t* tid2eid, + const float* bias_vl, + int64_t image_token_start, + cudaStream_t stream) { + cudaLaunchConfig_t config{}; + config.gridDim = (logits.size(0) + kWarpsPerBlock - 1) / kWarpsPerBlock; + config.blockDim = kThreadsPerBlock; + config.stream = stream; + + cudaLaunchAttribute attribute{}; + if constexpr (UsePDL) { + attribute.id = cudaLaunchAttributeProgrammaticStreamSerialization; + attribute.val.programmaticStreamSerializationAllowed = true; + config.attrs = &attribute; + config.numAttrs = 1; + } + + C10_CUDA_CHECK(cudaLaunchKernelEx( + &config, + topk_softplus_sqrt_kernel, + logits.data_ptr(), + topk_weights.data_ptr(), + topk_indices.data_ptr(), + static_cast(logits.size(0)), + static_cast(logits.size(1)), + static_cast(routed_scaling_factor), + correction_bias, + input_ids, + tid2eid, + bias_vl, + image_token_start)); +} + +void dispatch_pdl( + const at::Tensor& topk_weights, + const at::Tensor& topk_indices, + const at::Tensor& logits, + double routed_scaling_factor, + const float* correction_bias, + const int64_t* input_ids, + const int64_t* tid2eid, + const float* bias_vl, + int64_t image_token_start, + cudaStream_t stream) { + if (tid2eid == nullptr) { + if (logits.size(1) == 256) { + if (at::cuda::getCurrentDeviceProperties()->major >= 9) { + launch_topk_bias_softplus_sqrt<256, true>( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + bias_vl, + image_token_start, + stream); + } else { + launch_topk_bias_softplus_sqrt<256, false>( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + bias_vl, + image_token_start, + stream); + } + } else if (at::cuda::getCurrentDeviceProperties()->major >= 9) { + launch_topk_bias_softplus_sqrt<384, true>( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + bias_vl, + image_token_start, + stream); + } else { + launch_topk_bias_softplus_sqrt<384, false>( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + bias_vl, + image_token_start, + stream); + } + return; + } + + if (at::cuda::getCurrentDeviceProperties()->major >= 9) { + launch_topk_softplus_sqrt( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + tid2eid, + bias_vl, + image_token_start, + stream); + } else { + launch_topk_softplus_sqrt( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias, + input_ids, + tid2eid, + bias_vl, + image_token_start, + stream); + } +} + +void check_optional_vector( + const at::Tensor& logits, + const c10::optional& tensor, + at::ScalarType dtype, + int64_t length, + const char* name) { + if (!tensor.has_value()) return; + TORCH_CHECK( + tensor->is_cuda() && tensor->get_device() == logits.get_device() && tensor->scalar_type() == dtype && + tensor->dim() == 1 && tensor->size(0) == length && tensor->is_contiguous(), + name, + " must be a contiguous vector on the logits device"); +} + +} // namespace + +void topk_softplus_sqrt_cuda( + const at::Tensor& topk_weights, + const at::Tensor& topk_indices, + const at::Tensor& logits, + double routed_scaling_factor, + const c10::optional& correction_bias, + const c10::optional& input_ids, + const c10::optional& tid2eid, + const c10::optional& bias_vl, + int64_t image_token_start) { + TORCH_CHECK( + logits.is_cuda() && logits.scalar_type() == at::kFloat && logits.dim() == 2 && logits.is_contiguous(), + "logits must be contiguous [num_tokens, num_experts] float32 CUDA"); + TORCH_CHECK(logits.size(1) == 256 || logits.size(1) == 384, "DSV4 router requires 256 or 384 experts"); + TORCH_CHECK( + topk_weights.is_cuda() && topk_weights.get_device() == logits.get_device() && + topk_weights.scalar_type() == at::kFloat && topk_weights.is_contiguous() && + topk_weights.sizes() == at::IntArrayRef({logits.size(0), kTopK}), + "topk_weights must be contiguous [num_tokens, 6] float32 on the logits device"); + TORCH_CHECK( + topk_indices.is_cuda() && topk_indices.get_device() == logits.get_device() && topk_indices.is_contiguous() && + topk_indices.scalar_type() == at::kLong && + topk_indices.sizes() == at::IntArrayRef({logits.size(0), kTopK}), + "topk_indices must be contiguous [num_tokens, 6] int64 on the logits device"); + + check_optional_vector(logits, correction_bias, at::kFloat, logits.size(1), "correction_bias"); + check_optional_vector(logits, bias_vl, at::kFloat, logits.size(1), "bias_vl"); + if (input_ids.has_value()) { + TORCH_CHECK( + input_ids->is_cuda() && input_ids->get_device() == logits.get_device() && input_ids->dim() == 1 && + input_ids->size(0) == logits.size(0) && input_ids->is_contiguous() && input_ids->scalar_type() == at::kLong, + "input_ids must be contiguous [num_tokens] int64 on the logits device"); + } + if (tid2eid.has_value()) { + TORCH_CHECK(input_ids.has_value(), "input_ids is required for hash routing"); + TORCH_CHECK( + tid2eid->is_cuda() && tid2eid->get_device() == logits.get_device() && tid2eid->dim() == 2 && + tid2eid->size(1) == kTopK && tid2eid->is_contiguous() && tid2eid->scalar_type() == at::kLong, + "tid2eid must be contiguous [vocab_size, 6] int64 on the logits device"); + } else { + TORCH_CHECK(correction_bias.has_value(), "correction_bias is required for non-hash routing"); + } + if (bias_vl.has_value()) { + TORCH_CHECK(input_ids.has_value() && image_token_start > 0, "vision routing requires input_ids and image_token_start"); + } + if (logits.size(0) == 0) return; + + c10::cuda::CUDAGuard guard(logits.device()); + const auto stream = at::cuda::getCurrentCUDAStream(); + const float* correction_bias_ptr = + correction_bias.has_value() ? correction_bias->data_ptr() : nullptr; + const float* bias_vl_ptr = bias_vl.has_value() ? bias_vl->data_ptr() : nullptr; + dispatch_pdl( + topk_weights, + topk_indices, + logits, + routed_scaling_factor, + correction_bias_ptr, + input_ids.has_value() ? input_ids->data_ptr() : nullptr, + tid2eid.has_value() ? tid2eid->data_ptr() : nullptr, + bias_vl_ptr, + image_token_start, + stream); +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) { + module.def("topk_softplus_sqrt", &topk_softplus_sqrt_cuda, "DeepSeek-V4 fused top-6 sqrt-softplus router"); +} diff --git a/lightllm/models/deepseek_v4/triton_kernel/topk_softplus_sqrt.py b/lightllm/models/deepseek_v4/triton_kernel/topk_softplus_sqrt.py new file mode 100644 index 0000000000..a595760bcb --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/topk_softplus_sqrt.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the LightLLM project +# +# Wraps the CUDA router adapted from vLLM commit +# 36d2a5086bae12d0c0b311607373e4fc2d1036aa. + +import functools +import hashlib +import os + +import torch + + +@functools.lru_cache(maxsize=1) +def _load_cuda(): + from torch.utils.cpp_extension import load + + src = os.path.join(os.path.dirname(__file__), "csrc", "topk_softplus_sqrt.cu") + flags = ["-O3"] + with open(src, "rb") as source_file: + source = source_file.read() + capability = torch.cuda.get_device_capability() + cache_key = b"\0".join( + [ + source, + " ".join(flags).encode(), + torch.__version__.encode(), + str(torch.version.cuda).encode(), + f"sm{capability[0]}{capability[1]}".encode(), + os.environ.get("TORCH_CUDA_ARCH_LIST", "").encode(), + ] + ) + module_name = f"lightllm_dsv4_router_v1_{hashlib.sha256(cache_key).hexdigest()[:16]}" + return load( + name=module_name, + sources=[src], + extra_cuda_cflags=flags, + verbose=False, + ) + + +@torch.no_grad() +def topk_softplus_sqrt( + topk_weights: torch.Tensor, + topk_indices: torch.Tensor, + logits: torch.Tensor, + routed_scaling_factor: float, + correction_bias: torch.Tensor = None, + input_ids: torch.Tensor = None, + tid2eid: torch.Tensor = None, + bias_vl: torch.Tensor = None, + image_token_start: int = 0, +) -> None: + _load_cuda().topk_softplus_sqrt( + topk_weights, + topk_indices, + logits, + float(routed_scaling_factor), + correction_bias, + input_ids, + tid2eid, + bias_vl, + int(image_token_start), + ) + return diff --git a/lightllm/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py index c5b581b139..4c8225b000 100644 --- a/lightllm/models/deepseek_v4/workspace.py +++ b/lightllm/models/deepseek_v4/workspace.py @@ -7,13 +7,14 @@ def __init__(self, model): self.token_capacity = int(model.batch_max_tokens) self.sliding_window = int(model.config["sliding_window"]) args = get_env_start_args() - self.swa_capacity = self.sliding_window + self.vision_swa_width = self.sliding_window + int(model.config.get("vision_max_n_token", 0)) + self.swa_capacity = self.vision_swa_width if args.mtp_mode == "dspark": dspark_width = self.sliding_window + int(args.mtp_step) # FlashMLA sparse decode requires the physical top-k width to be block aligned. # 128 covers both supported padded Q-head configurations; swa_lengths keeps # the actual number of visible history + draft-block entries. - self.swa_capacity = ((dspark_width + 127) // 128) * 128 + self.swa_capacity = max(self.swa_capacity, ((dspark_width + 127) // 128) * 128) self.index_topk = int(model.config["index_topk"]) self.c128_cap = self.compress_cap(model.max_seq_length, 128) overlap = args.enable_decode_microbatch_overlap or args.enable_prefill_microbatch_overlap diff --git a/lightllm/models/gemma4/tokenizer.py b/lightllm/models/gemma4/tokenizer.py index 5a675856f2..760203b8da 100644 --- a/lightllm/models/gemma4/tokenizer.py +++ b/lightllm/models/gemma4/tokenizer.py @@ -78,7 +78,9 @@ def encode(self, prompt, multimodal_params: MultimodalParams = None, add_special if not input_ids or input_ids[-1] != self.boi_token_index: input_ids.append(self.boi_token_index) img.start_idx = len(input_ids) + img.block_start_idx = img.start_idx input_ids.extend(range(img.token_id, img.token_id + img.token_num)) + img.block_end_idx = len(input_ids) input_ids.append(self.eoi_token_index) if image_end < len(origin_ids) and origin_ids[image_end] == self.eoi_token_index: diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index c02b5a07e8..8f8f3dede4 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -388,6 +388,14 @@ async def generate( await self._log_req_header(request_headers, group_request_id) # encode prompt_ids = await self._encode(prompt, multimodal_params, sampling_params) + for image in multimodal_params.images: + if image.block_start_idx is not None: + block_token_num = image.block_end_idx - image.block_start_idx + if block_token_num > self.args.batch_max_tokens: + raise ValueError( + f"image prefill block token count {block_token_num} exceeds " + f"batch_max_tokens={self.args.batch_max_tokens}; increase --batch_max_tokens" + ) self._log_stage_timing( group_request_id, start_time, diff --git a/lightllm/server/multimodal_params.py b/lightllm/server/multimodal_params.py index ed3535a69f..8787e3ca8e 100644 --- a/lightllm/server/multimodal_params.py +++ b/lightllm/server/multimodal_params.py @@ -1,4 +1,5 @@ """Multimodal parameters for text generation.""" + import asyncio import os import librosa @@ -127,6 +128,9 @@ def __init__(self, **kwargs): # the start index of the image in the input_ids # used for mrope position id calculation self.start_idx = None + # optional half-open input range that must be processed in one prefill step + self.block_start_idx = None + self.block_end_idx = None self.grid_thwd = None self.image_w = 0 self.image_h = 0 @@ -217,6 +221,8 @@ def to_dict(self): ret["token_num"] = self.token_num ret["grid_thwd"] = self.grid_thwd ret["start_idx"] = self.start_idx + ret["block_start_idx"] = self.block_start_idx + ret["block_end_idx"] = self.block_end_idx ret["md5"] = self.md5 return ret diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index f0bb77aa98..ee99cf2720 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -719,6 +719,11 @@ def _init_all_state(self): self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() self.multimodal_params = self.multimodal_params.to_dict() + self.image_block_spans = [ + (image["block_start_idx"], image["block_end_idx"]) + for image in self.multimodal_params["images"] + if image["block_start_idx"] is not None + ] self.shared_kv_node: Union[TreeNode, LinearAttPagedTreeNode] = None self.finish_status = FinishStatus() @@ -740,6 +745,16 @@ def _match_radix_cache(self): input_token_ids = self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] key = torch.tensor(input_token_ids, dtype=torch.int64, device="cpu") key = key[0 : len(key) - 1] # 最后一个不需要,因为需要一个额外的token,让其在prefill的时候输出下一个token的值 + # DSV4 prompt cache may reclaim earlier SWA pages, so an image-internal hit must be recomputed. + if g_infer_context.is_deepseek_v4 and self.image_block_spans: + while True: + _, matched_len, _ = g_infer_context.radix_cache.match_prefix(key, update_refs=False) + for image_start, image_end in self.image_block_spans: + if image_start < matched_len < image_end: + key = key[:image_start] + break + else: + break share_node, kv_len, value_tensor = g_infer_context.radix_cache.match_prefix(key, update_refs=True) if share_node is not None: self.shared_kv_node = share_node @@ -922,10 +937,17 @@ def get_cur_total_len(self): def get_input_token_ids(self): return self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] - def get_chuncked_input_token_ids(self): + def _get_chunked_input_end(self): chunked_start = self.cur_kv_len chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) - return self.shm_req.shm_prompt_ids.arr[0:chunked_end] + for image_start, image_end in self.image_block_spans: + if image_start < chunked_end < image_end: + chunked_end = image_start if chunked_start < image_start else image_end + break + return chunked_end + + def get_chuncked_input_token_ids(self): + return self.shm_req.shm_prompt_ids.arr[0 : self._get_chunked_input_end()] def get_chuncked_input_token_ids_for_linear_att(self): big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num @@ -943,9 +965,7 @@ def get_chuncked_input_token_ids_for_linear_att(self): return self.shm_req.shm_prompt_ids.arr[0:end] def get_chuncked_input_token_len(self): - chunked_start = self.cur_kv_len - chunked_end = min(self.get_cur_total_len(), chunked_start + self.args.chunked_prefill_size) - return chunked_end + return self._get_chunked_input_end() def get_chuncked_input_token_len_for_linear_att(self): big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num @@ -1066,9 +1086,13 @@ def get_dsv4_recover_need_page_and_slot_num(self) -> Tuple[int, int, int]: # C4/C128 accumulate across recovery chunks; only SWA is evicted chunk by chunk. req_manager: DeepseekV4ReqManager = g_infer_context.req_manager prompt_cache_page_size = req_manager.get_prompt_cache_page_size() + max_prefill_token_num = max( + self.args.chunked_prefill_size, + max((image_end - image_start for image_start, image_end in self.image_block_spans), default=0), + ) peak_token_num = min( self.get_cur_total_len(), - self.args.chunked_prefill_size + int(req_manager.sliding_window) + 2 * prompt_cache_page_size, + max_prefill_token_num + int(req_manager.sliding_window) + 2 * prompt_cache_page_size, ) swa_page_num = (peak_token_num + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size return swa_page_num, c4_page_num, c128_slot_num diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py index eef20b9a57..f36d72af07 100644 --- a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py @@ -243,6 +243,14 @@ def store_completed_prefill_pages( session.closing = True self._try_release_dsv4_session(session) + @staticmethod + def _get_image_safe_load_end(req: InferReq, loaded_start: int, load_end: int, page_size: int) -> int: + """Move an image-internal CPU resume point before that image.""" + for image_start, image_end in reversed(req.image_block_spans): + if image_start < load_end < image_end: + load_end = image_start // page_size * page_size + return load_end if load_end > loaded_start else 0 + def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): idle_token_num = g_infer_context.get_can_alloc_token_num() is_master_in_dp = self.backend.is_master_in_dp @@ -282,6 +290,7 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): c4_capacity, c128_capacity, ) + loadable_end = self._get_image_safe_load_end(req, gpu_kv_len, loadable_end, layout.token_page_size) if loadable_end != 0: token_num = loadable_end - gpu_kv_len full_need = token_num @@ -305,6 +314,9 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): int(mem_manager.c4_page_allocator.can_use_mem_size) if mem_manager.n_c4 else 0, int(mem_manager.c128_allocator.can_use_mem_size) if mem_manager.n_c128 else 0, ) + loadable_end = self._get_image_safe_load_end( + req, gpu_kv_len, loadable_end, layout.token_page_size + ) if loadable_end != 0: loaded_end = loadable_end token_num = loaded_end - gpu_kv_len diff --git a/lightllm/server/tokenizer.py b/lightllm/server/tokenizer.py index 1e7522ad66..b6d0048b88 100644 --- a/lightllm/server/tokenizer.py +++ b/lightllm/server/tokenizer.py @@ -99,7 +99,7 @@ def get_tokenizer( from ..models.deepseek_v4.model import DeepSeekV4Tokenizer logger.info("Using DeepSeek-V4 tokenizer mode with Python-based chat template encoding.") - return DeepSeekV4Tokenizer(tokenizer, tokenizer_name) + return DeepSeekV4Tokenizer(tokenizer, tokenizer_name, model_cfg) if model_cfg["architectures"][0] == "TarsierForConditionalGeneration": from ..models.qwen2_vl.vision_process import Qwen2VLImageProcessor diff --git a/lightllm/server/visualserver/model_infer/model_rpc.py b/lightllm/server/visualserver/model_infer/model_rpc.py index 68e0a97ca1..e376d6534f 100644 --- a/lightllm/server/visualserver/model_infer/model_rpc.py +++ b/lightllm/server/visualserver/model_infer/model_rpc.py @@ -14,6 +14,7 @@ from lightllm.models.llava.llava_visual import LlavaVisionModel from lightllm.models.gemma3.gemma3_visual import Gemma3VisionModel from lightllm.models.gemma4.gemma4_visual import Gemma4VisionModel +from lightllm.models.deepseek_v4.deepseek_v4_visual import DeepseekV4VisionModel from lightllm.models.vit.model import VisionTransformer from lightllm.server.multimodal_params import MultimodalParams, ImageItem from lightllm.models.qwen2_vl.qwen2_visual import Qwen2VisionTransformerPretrainedModel @@ -91,6 +92,8 @@ def exposed_init_model(self, kvargs): self.model = ( Qwen3VisionTransformerPretrainedModel(kvargs, **model_cfg["vision_config"]).eval().bfloat16() ) + elif self.model_type == "deepseek_v4" and model_cfg.get("vision_n_layers", 0) > 0: + self.model = DeepseekV4VisionModel(kvargs, **model_cfg).eval() elif model_cfg["architectures"][0] == "TarsierForConditionalGeneration": self.model = TarsierVisionTransformerPretrainedModel(**model_cfg).eval().bfloat16() elif self.model_type == "llava": diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 23a38c368f..66ba9503ce 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -472,6 +472,8 @@ def has_vision_module(model_path: str) -> bool: # Qwen3VisionTransformerPretrainedModel model_cfg["vision_config"] return True + elif model_type == "deepseek_v4": + return model_cfg.get("vision_n_layers", 0) > 0 elif model_cfg["architectures"][0] == "TarsierForConditionalGeneration": # TarsierVisionTransformerPretrainedModel return True diff --git a/unit_tests/models/deepseek_v4/test_vision.py b/unit_tests/models/deepseek_v4/test_vision.py new file mode 100644 index 0000000000..96343b614c --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_vision.py @@ -0,0 +1,294 @@ +import importlib.util +import io +import json +import os +import sys +from types import SimpleNamespace + +import numpy as np +import pytest +import torch +from PIL import Image +from torch import nn + +from lightllm.models.deepseek_v4 import deepseek_v4_visual as visual +from lightllm.models.deepseek_v4 import image_processor as ours + + +REFERENCE_DIR = "/mtc/models/DeepSeek-V4-Flash-Vision-Exp/inference" +REFERENCE_PROCESSOR_PATH = os.path.join(REFERENCE_DIR, "image_processor.py") +REFERENCE_VISION_PATH = os.path.join(REFERENCE_DIR, "vision.py") + +PATCH_SIZE = 14 +DOWNSAMPLE_RATIO = 3 +MAX_N_TOKEN = 384 +MIN_PIXELS = 147456 +MAX_WH_RATIO = 8 + +REFERENCE_ARGS = SimpleNamespace( + vision_patch_size=PATCH_SIZE, + vision_downsample_ratio=DOWNSAMPLE_RATIO, + vision_max_n_token=MAX_N_TOKEN, + vision_min_pixels=MIN_PIXELS, + vision_max_wh_ratio=MAX_WH_RATIO, +) + + +def _load_reference_processor(): + spec = importlib.util.spec_from_file_location("deepseek_v4_reference_image_processor", REFERENCE_PROCESSOR_PATH) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def _load_reference_vision(): + spec = importlib.util.spec_from_file_location("deepseek_v4_reference_vision", REFERENCE_VISION_PATH) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def reference_processor(): + if not os.path.exists(REFERENCE_PROCESSOR_PATH): + pytest.skip("DeepSeek-V4 vision reference processor is not available") + return _load_reference_processor() + + +@pytest.fixture(scope="module") +def reference_vision(): + if not os.path.exists(REFERENCE_VISION_PATH): + pytest.skip("DeepSeek-V4 reference vision tower is not available") + return _load_reference_vision() + + +@pytest.mark.parametrize( + "height,width", + [(14, 14), (15, 30), (41, 512), (100, 100), (410, 756), (756, 42)], +) +def test_grid_tokens_matches_reference(reference_processor, height, width): + assert ours.grid_tokens(height, width, PATCH_SIZE, DOWNSAMPLE_RATIO) == reference_processor.grid_tokens( + height, width, PATCH_SIZE, DOWNSAMPLE_RATIO + ) + + +@pytest.mark.parametrize("start_pos", range(8)) +@pytest.mark.parametrize("n_llm_h,n_llm_w", [(1, 1), (2, 3), (3, 2), (4, 5), (7, 8)]) +def test_image_block_and_permutation_match_reference(reference_processor, n_llm_h, n_llm_w, start_pos): + expected_types, expected_perm = reference_processor.build_image_block(n_llm_h, n_llm_w, start_pos) + types, perm = ours.build_image_block(n_llm_h, n_llm_w, start_pos) + assert torch.equal(types, expected_types) + assert torch.equal(perm, expected_perm) + + +@pytest.mark.parametrize( + "width,height", + [(800, 600), (100, 2000), (2000, 100), (50, 50), (13, 7), (384, 384)], +) +def test_image_geometry_and_pixels_match_reference(reference_processor, width, height): + rng = np.random.default_rng(width * 10000 + height) + image = Image.fromarray(rng.integers(0, 256, (height, width, 3), dtype=np.uint8)) + image_bytes = io.BytesIO() + image.save(image_bytes, format="PNG") + + expected = reference_processor.load_image({"data": image_bytes.getvalue()}, REFERENCE_ARGS) + actual = ours.load_image(image, REFERENCE_ARGS) + assert actual[1:] == expected[1:] + assert torch.equal(actual[0], expected[0]) + + n_vit_h, n_vit_w, n_llm_h, n_llm_w, token_num = ours.get_image_grid( + height, + width, + PATCH_SIZE, + DOWNSAMPLE_RATIO, + MAX_N_TOKEN, + MIN_PIXELS, + MAX_WH_RATIO, + ) + assert (n_vit_h, n_vit_w, n_llm_h, n_llm_w) == actual[1:] + reference_types, _ = reference_processor.build_image_block(n_llm_h, n_llm_w, start_pos=0) + assert token_num == len(reference_types) + assert token_num <= MAX_N_TOKEN + + +def _tiny_model(): + return visual.DeepseekV4VisionModel( + {"weight_dir": "/unused"}, + hidden_size=8, + vision_n_layers=1, + vision_dim=8, + vision_n_heads=2, + vision_inter_dim=12, + vision_patch_size=2, + vision_rope_theta=10000.0, + vision_downsample_ratio=1, + vision_max_n_token=64, + vision_min_pixels=0, + vision_max_wh_ratio=8, + ) + + +def _tiny_vision_args(): + return SimpleNamespace( + hidden_size=12, + dim=12, + vision_n_layers=2, + vision_dim=16, + vision_n_heads=4, + vision_inter_dim=24, + vision_patch_size=2, + vision_rope_theta=10000.0, + vision_downsample_ratio=2, + ) + + +@pytest.mark.parametrize("n_vit_h,n_vit_w", [(4, 4), (5, 3)]) +def test_vit_matches_reference(reference_vision, n_vit_h, n_vit_w): + torch.manual_seed(0) + args = _tiny_vision_args() + model = visual.ViT(args).eval() + previous_dtype = torch.get_default_dtype() + torch.set_default_dtype(torch.bfloat16) + try: + reference = reference_vision.ViT(args).eval() + finally: + torch.set_default_dtype(previous_dtype) + reference.load_state_dict(model.state_dict()) + patches = torch.randn( + n_vit_h * n_vit_w, + 3, + args.vision_patch_size, + args.vision_patch_size, + dtype=torch.bfloat16, + ) + with torch.inference_mode(): + actual = model(patches, n_vit_h, n_vit_w) + expected = reference(patches, n_vit_h, n_vit_w) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.mark.parametrize("n_vit_h,n_vit_w", [(4, 4), (5, 3)]) +def test_aligner_matches_reference(reference_vision, n_vit_h, n_vit_w): + torch.manual_seed(0) + args = _tiny_vision_args() + model = visual.Aligner(args).eval() + previous_dtype = torch.get_default_dtype() + torch.set_default_dtype(torch.bfloat16) + try: + reference = reference_vision.Aligner(args).eval() + finally: + torch.set_default_dtype(previous_dtype) + reference.load_state_dict(model.state_dict()) + hidden = torch.randn(n_vit_h * n_vit_w, args.vision_dim, dtype=torch.bfloat16) + with torch.inference_mode(): + actual = model(hidden, n_vit_h, n_vit_w) + expected = reference(hidden, n_vit_h, n_vit_w) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_visual_parameter_dtypes(): + model = _tiny_model() + for module in model.modules(): + if isinstance(module, nn.Linear): + assert module.weight.dtype == torch.bfloat16 + if module.bias is not None: + assert module.bias.dtype == torch.bfloat16 + elif isinstance(module, visual.RMSNorm): + assert module.weight.dtype == torch.float32 + assert model.image_start.dtype == torch.bfloat16 + assert model.image_pad.dtype == torch.bfloat16 + + +def test_load_model_reads_only_indexed_visual_weights(tmp_path, monkeypatch): + model = _tiny_model() + expected = { + name: torch.full_like(tensor, index + 1) for index, (name, tensor) in enumerate(model.state_dict().items()) + } + weight_map = {name: "visual.safetensors" for name in expected} + weight_map["layers.0.self_attn.weight"] = "language.safetensors" + (tmp_path / "model.safetensors.index.json").write_text(json.dumps({"weight_map": weight_map})) + + requested = [] + opened = [] + + class FakeSafetensors: + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_value, traceback): + return False + + def get_tensor(self, key): + requested.append(key) + return expected[key] + + def fake_safe_open(path, framework, device): + opened.append(os.path.basename(path)) + assert framework == "pt" + assert device == "cpu" + return FakeSafetensors() + + monkeypatch.setattr(visual, "safe_open", fake_safe_open) + model.load_model(str(tmp_path)) + + assert opened == ["visual.safetensors"] + assert set(requested) == set(expected) + for name, tensor in model.state_dict().items(): + assert torch.equal(tensor, expected[name]) + + +def test_encode_returns_canonical_cache_block(monkeypatch): + model = _tiny_model() + with torch.no_grad(): + model.image_start.fill_(10) + model.image_pad.fill_(20) + model.image_newline.fill_(30) + model.image_end.fill_(40) + + class FakeVision(nn.Module): + def forward(self, patches, n_h, n_w): + return torch.zeros((n_h * n_w, 1), dtype=torch.bfloat16) + + class FakeAligner(nn.Module): + def forward(self, x, n_h, n_w): + values = torch.arange(n_h * n_w, dtype=torch.bfloat16) + return values[:, None].expand(-1, model.config.hidden_size) + + model.vision = FakeVision() + model.aligner = FakeAligner() + + image_bytes = io.BytesIO() + Image.new("RGB", (6, 4), color=(1, 2, 3)).save(image_bytes, format="PNG") + shm_names = [] + monkeypatch.setattr(visual, "get_shm_name_data", lambda uuid: f"{uuid}-data") + + def fake_read_shm(name): + shm_names.append(name) + return image_bytes.getvalue() + + monkeypatch.setattr(visual, "read_shm", fake_read_shm) + embeds, uuids, valid_ids = model.encode([SimpleNamespace(uuid="image-1")]) + + types, perm = ours.build_canonical_image_block(2, 3) + sentinel_table = torch.stack( + [ + model.image_start, + model.image_pad, + model.image_pad, + model.image_newline, + model.image_end, + ] + ) + expected = sentinel_table[types] + row_major = torch.arange(6, dtype=torch.bfloat16)[:, None].expand(-1, model.config.hidden_size) + expected[types == ours.IMAGE] = row_major[perm] + + assert shm_names == ["image-1-data"] + assert uuids == ["image-1"] + assert valid_ids == [[0, len(types)]] + assert torch.equal(embeds, expected) + assert torch.equal(types[:3], torch.full((3,), ours.IMAGE_PAD, dtype=torch.int64)) + assert types[3] == ours.IMAGE_START + assert types[-1] == ours.IMAGE_END diff --git a/unit_tests/models/deepseek_v4/test_vision_integration.py b/unit_tests/models/deepseek_v4/test_vision_integration.py new file mode 100644 index 0000000000..5a91049ec0 --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_vision_integration.py @@ -0,0 +1,408 @@ +from types import SimpleNamespace + +import pytest +import torch + + +VISION_CONFIG = { + "vision_n_layers": 32, + "vision_patch_size": 14, + "vision_downsample_ratio": 3, + "vision_max_n_token": 384, + "vision_min_pixels": 147456, + "vision_max_wh_ratio": 8, +} + + +class _Tokenizer: + def convert_tokens_to_ids(self, token): + assert token == "<|deepseek_image|>" + return 99 + + def encode(self, prompt, **kwargs): + return prompt + + +@pytest.mark.parametrize("prefix_len", range(8)) +def test_tokenizer_uses_position_dependent_canonical_suffix(prefix_len): + from lightllm.models.deepseek_v4.model import DeepSeekV4Tokenizer + + tokenizer = DeepSeekV4Tokenizer(_Tokenizer(), "/unused", VISION_CONFIG) + image = SimpleNamespace( + token_id=100_000, + token_num=13, + start_idx=None, + block_start_idx=None, + block_end_idx=None, + ) + params = SimpleNamespace(images=[image]) + prompt = [7] * prefix_len + [99, 8] + + result = tokenizer.encode(prompt, params) + + cache_skip = prefix_len % 4 + assert result == [7] * prefix_len + list(range(image.token_id + cache_skip, image.token_id + 13)) + [8] + assert image.block_start_idx == prefix_len + assert image.block_end_idx == prefix_len + image.token_num - cache_skip + assert image.start_idx == prefix_len + 3 - cache_skip + assert image.start_idx % 4 == 3 + assert result[image.block_start_idx : image.block_end_idx] == list( + range(image.token_id + cache_skip, image.token_id + image.token_num) + ) + + +def test_tokenizer_keeps_dsv4_on_one_dimensional_positions(): + from lightllm.models.deepseek_v4.model import DeepSeekV4Tokenizer + + tokenizer = DeepSeekV4Tokenizer(_Tokenizer(), "/unused", VISION_CONFIG) + image = SimpleNamespace(image_h=64, image_w=64, grid_thwd=None) + + assert tokenizer.get_image_token_length(image) > 0 + assert image.grid_thwd is None + + +def test_chunk_boundary_moves_before_complete_image_block(): + from lightllm.server.router.model_infer.infer_batch import InferReq + + req = InferReq.__new__(InferReq) + req.args = SimpleNamespace(chunked_prefill_size=8192) + req.cur_kv_len = 0 + req.cur_output_len = 0 + req.shm_req = SimpleNamespace(input_len=10_000) + req.image_block_spans = [(8100, 8400)] + + assert req._get_chunked_input_end() == 8100 + + req.cur_kv_len = 8100 + assert req._get_chunked_input_end() == 10_000 + + +def test_chunk_boundary_between_adjacent_images_never_splits_the_previous_image(): + from lightllm.server.router.model_infer.infer_batch import InferReq + + req = InferReq.__new__(InferReq) + req.args = SimpleNamespace(chunked_prefill_size=640) + req.cur_kv_len = 0 + req.cur_output_len = 0 + req.shm_req = SimpleNamespace(input_len=768) + req.image_block_spans = [(0, 384), (384, 768)] + + assert req._get_chunked_input_end() == 384 + + req.cur_kv_len = 384 + assert req._get_chunked_input_end() == 768 + + +def test_atomic_image_block_can_exceed_chunk_size(): + from lightllm.server.router.model_infer.infer_batch import InferReq + + req = InferReq.__new__(InferReq) + req.args = SimpleNamespace(chunked_prefill_size=256) + req.cur_kv_len = 0 + req.cur_output_len = 0 + req.shm_req = SimpleNamespace(input_len=1000) + req.image_block_spans = [(100, 800)] + + assert req._get_chunked_input_end() == 100 + + req.cur_kv_len = 100 + assert req._get_chunked_input_end() == 800 + + +def test_radix_match_retries_when_page_alignment_lands_in_earlier_image(monkeypatch): + from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context + + class FakeRadixCache: + def __init__(self): + self.calls = [] + + def match_prefix(self, key, update_refs=False): + self.calls.append((len(key), update_refs)) + matched_len = min(len(key) // 256 * 256, 768) + if update_refs: + node = SimpleNamespace(node_prefix_total_len=matched_len) + return node, matched_len, torch.arange(matched_len) + return None, matched_len, None + + radix_cache = FakeRadixCache() + monkeypatch.setattr(g_infer_context, "is_linear_att_mixed_model", False) + monkeypatch.setattr(g_infer_context, "is_deepseek_v4", True) + monkeypatch.setattr(g_infer_context, "radix_cache", radix_cache) + monkeypatch.setattr( + g_infer_context, + "req_manager", + SimpleNamespace(req_to_token_indexs=torch.empty((1, 1025), dtype=torch.int64)), + ) + + req = InferReq.__new__(InferReq) + req.sampling_param = SimpleNamespace(disable_prompt_cache=False) + req.cur_kv_len = 0 + req.cur_output_len = 0 + req.req_idx = 0 + req.image_block_spans = [(450, 600), (700, 900)] + req.shared_kv_node = None + req.shm_req = SimpleNamespace( + input_len=1025, + shm_prompt_ids=SimpleNamespace(arr=list(range(1025))), + prompt_cache_len=0, + shm_cur_kv_len=0, + ) + + req._match_radix_cache() + + assert radix_cache.calls == [ + (1024, False), + (700, False), + (450, False), + (450, True), + ] + assert req.cur_kv_len == 256 + assert req.shm_req.prompt_cache_len == 256 + + +def test_recover_swa_budget_includes_atomic_image_block(monkeypatch): + from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context + + monkeypatch.setattr( + g_infer_context, + "req_manager", + SimpleNamespace(sliding_window=128, get_prompt_cache_page_size=lambda: 256), + ) + + req = InferReq.__new__(InferReq) + req.args = SimpleNamespace(disable_chunked_prefill=False, chunked_prefill_size=128) + req.cur_kv_len = 0 + req.cur_output_len = 0 + req.shm_req = SimpleNamespace(input_len=2000) + req.image_block_spans = [(500, 884)] + req.dsv4_swa_page_size = 128 + req.dsv4_c4_page_size = 64 + req.dsv4_has_c128 = False + + assert req.get_dsv4_recover_need_page_and_slot_num() == (8, 8, 0) + + +@pytest.mark.parametrize( + ("loaded_start", "load_end", "spans", "expected"), + [ + (0, 4096, [], 4096), + (0, 2048, [(2048, 2500)], 2048), + (0, 4096, [(3500, 4096)], 4096), + (0, 4096, [(3500, 4500)], 2048), + (0, 4096, [(1000, 2500), (3500, 4500)], 0), + (2304, 4096, [(3000, 4500)], 0), + ], +) +def test_cpu_cache_load_end_never_splits_an_image(loaded_start, load_end, spans, expected): + from lightllm.server.router.model_infer.mode_backend.dsv4_multi_level_kv_cache import ( + Dsv4MultiLevelKvCacheModule, + ) + + req = SimpleNamespace(image_block_spans=spans) + + assert Dsv4MultiLevelKvCacheModule._get_image_safe_load_end(req, loaded_start, load_end, 2048) == expected + + +def test_cpu_cache_rechecks_image_boundary_after_capacity_changes(monkeypatch): + from lightllm.server.router.model_infer.mode_backend import ( + dsv4_multi_level_kv_cache as cache_module, + ) + + capacity_results = iter([6144, 4096]) + capacity_calls = [] + prepare_calls = [] + loaded_pages = [] + finish_calls = [] + + def get_loadable_cpu_cache_end(*args): + capacity_calls.append(args) + return next(capacity_results) + + def prepare_cpu_cache_load(*, token_num, loaded_end): + prepare_calls.append((token_num, loaded_end)) + return SimpleNamespace(mem_indexes=torch.arange(token_num, dtype=torch.int32)) + + def load_cpu_cache_pages(*, page_indexes, **kwargs): + loaded_pages.append(page_indexes.tolist()) + + req_to_token_indexs = torch.full((1, 6144), -1, dtype=torch.int32) + req_manager = SimpleNamespace( + req_to_token_indexs=req_to_token_indexs, + finish_cpu_cache_load=lambda req_idx, loaded_end: finish_calls.append((req_idx, loaded_end)), + ) + mem_manager = SimpleNamespace( + cpu_cache_layout=SimpleNamespace(token_page_size=2048), + n_c4=False, + n_c128=False, + allocator=SimpleNamespace(can_use_mem_size=8192), + swa_page_allocator=SimpleNamespace(can_use_mem_size=2), + get_loadable_cpu_cache_end=get_loadable_cpu_cache_end, + prepare_cpu_cache_load=prepare_cpu_cache_load, + operator=SimpleNamespace(load_cpu_cache_pages=load_cpu_cache_pages), + commit_cpu_cache_load_plan=lambda plan: None, + ) + module = object.__new__(cache_module.Dsv4MultiLevelKvCacheModule) + module.backend = SimpleNamespace( + is_master_in_dp=True, + radix_cache=None, + model=SimpleNamespace(mem_manager=mem_manager, req_manager=req_manager), + ) + module.cpu_cache_client = object() + module.init_sync_group = None + module._dsv4_store_sessions = {} + + req = SimpleNamespace( + req_id=7, + req_idx=0, + cur_kv_len=0, + image_block_spans=[(3500, 4500)], + shm_req=SimpleNamespace( + cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [10, 11, 12]), + token_hash_page_len_list=SimpleNamespace(get_all=lambda: [2048, 4096, 6144]), + disk_prompt_cache_len=2048, + cpu_prompt_cache_len=0, + shm_cur_kv_len=0, + ), + ) + + real_tensor = torch.tensor + monkeypatch.setattr( + cache_module.torch, + "tensor", + lambda data, **kwargs: real_tensor(data, **{key: value for key, value in kwargs.items() if key != "device"}), + ) + monkeypatch.setattr(cache_module.torch.cuda, "Event", lambda: SimpleNamespace(record=lambda: None)) + monkeypatch.setattr(cache_module.dist, "barrier", lambda group: None) + monkeypatch.setattr(cache_module.g_infer_context, "get_can_alloc_token_num", lambda: 8192) + monkeypatch.setattr( + cache_module.g_infer_context, + "get_can_alloc_dsv4_page_and_slot_num", + lambda: (2, 0, 0), + ) + + module.load_cpu_cache_to_reqs([req]) + + assert len(capacity_calls) == 2 + assert prepare_calls == [(2048, 2048)] + assert loaded_pages == [[10]] + assert finish_calls == [(0, 2048)] + assert req.cur_kv_len == 2048 + assert req.shm_req.shm_cur_kv_len == 2048 + assert req.shm_req.cpu_prompt_cache_len == 2048 + assert req.shm_req.disk_prompt_cache_len == 0 + assert module._dsv4_store_sessions[7].leased_pages == [10, 11, 12] + + +def test_vision_model_allows_cpu_cache(): + from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + + model = SimpleNamespace( + load_way="HF", + tp_world_size_=1, + config={ + "num_attention_heads": 8, + "o_groups": 8, + "index_n_heads": 8, + "vision_n_layers": 1, + }, + args=SimpleNamespace(enable_cpu_cache=True), + ) + + DeepseekV4TpPartModel._verify_params(model) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("is_hash", [False, True]) +def test_router_uses_bias_vl_only_for_image_tokens(is_hash): + from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer + + router = DeepseekV4TransformerLayerInfer.__new__(DeepseekV4TransformerLayerInfer) + router.has_vision = True + router.is_hash = is_hash + router.vocab_size = 32 + router.num_experts_per_tok = 6 + router.routed_scaling_factor = 1.5 + router.alloc_tensor = torch.empty + + torch.manual_seed(0) + logits = torch.randn((2, 256), dtype=torch.float32, device="cuda") + hash_table = torch.zeros((router.vocab_size, 6), dtype=torch.int64, device="cuda") + hash_table[2] = torch.tensor([1, 4, 7, 10, 13, 16], device="cuda") + text_bias = torch.zeros(256, dtype=torch.float32, device="cuda") + text_bias[20:26] = torch.tensor([60.0, 55.0, 50.0, 45.0, 40.0, 35.0], device="cuda") + vision_bias = torch.zeros(256, dtype=torch.float32, device="cuda") + vision_bias[200:206] = torch.tensor([60.0, 55.0, 50.0, 45.0, 40.0, 35.0], device="cuda") + layer_weight = SimpleNamespace( + gate_tid2eid_=SimpleNamespace(weight=hash_table), + gate_bias_=SimpleNamespace(weight=text_bias), + gate_bias_vl_=SimpleNamespace(weight=vision_bias), + ) + infer_state = SimpleNamespace(is_prefill=True, input_ids=torch.tensor([2, 100_000], device="cuda")) + + weights, indices = router._select_experts(logits, infer_state, layer_weight) + + scores = torch.sqrt(torch.nn.functional.softplus(logits)) + text_indices = hash_table[infer_state.input_ids[0]] if is_hash else (scores[0] + text_bias).topk(6).indices + image_indices = (scores[1] + vision_bias).topk(6).indices + expected_indices = torch.stack((text_indices, image_indices)) + torch.testing.assert_close(indices, expected_indices) + expected = scores.gather(1, expected_indices) + expected = expected / expected.sum(dim=-1, keepdim=True) * 1.5 + torch.testing.assert_close(weights, expected, rtol=2e-5, atol=1e-6) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_build_image_visibility_scatter(): + from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_image_visibility + + image_spans = torch.tensor([[0, 1, 4], [0, 7, 3], [1, 6, 3]], dtype=torch.int32, device="cuda") + b_q_start_loc = torch.tensor([0, 6], dtype=torch.int32, device="cuda") + b_ready_cache_len = torch.tensor([2, 5], dtype=torch.int32, device="cuda") + b_q_seq_len = torch.tensor([6, 5], dtype=torch.int32, device="cuda") + image_left = torch.zeros(11, dtype=torch.int32, device="cuda") + image_right = torch.zeros(11, dtype=torch.int32, device="cuda") + + build_image_visibility( + image_spans, + b_q_start_loc, + b_ready_cache_len, + b_q_seq_len, + image_left, + image_right, + ) + + assert image_left.cpu().tolist() == [1, 2, 3, 0, 0, 0, 0, 0, 1, 2, 0] + assert image_right.cpu().tolist() == [2, 1, 0, 0, 0, 2, 0, 2, 1, 0, 0] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_swa_index_adds_bidirectional_image_visibility(): + from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_swa_index + + positions = torch.tensor([5, 8, 10], dtype=torch.int32, device="cuda") + req_idx = torch.zeros(3, dtype=torch.int32, device="cuda") + req_to_token = torch.arange(20, dtype=torch.int32, device="cuda").unsqueeze(0) + full_to_swa = torch.arange(100, 120, dtype=torch.int32, device="cuda") + output = torch.empty((3, 8), dtype=torch.int32, device="cuda") + lengths = torch.empty(3, dtype=torch.int32, device="cuda") + image_left = torch.tensor([0, 2, 4], dtype=torch.int32, device="cuda") + image_right = torch.tensor([0, 2, 0], dtype=torch.int32, device="cuda") + + build_swa_index( + req_idx, + positions, + req_to_token, + full_to_swa, + output, + lengths, + window=4, + image_left=image_left, + image_right=image_right, + ) + + assert output.cpu().tolist() == [ + [105, 104, 103, 102, -1, -1, -1, -1], + [105, 106, 107, 108, 109, 110, -1, -1], + [106, 107, 108, 109, 110, -1, -1, -1], + ] + assert lengths.cpu().tolist() == [4, 6, 5] diff --git a/unit_tests/models/gemma4/test_tokenizer.py b/unit_tests/models/gemma4/test_tokenizer.py new file mode 100644 index 0000000000..229f74039d --- /dev/null +++ b/unit_tests/models/gemma4/test_tokenizer.py @@ -0,0 +1,38 @@ +from types import SimpleNamespace + +from lightllm.models.gemma4.tokenizer import Gemma4Tokenizer +from lightllm.server.multimodal_params import ImageItem + + +class _Tokenizer: + bos_token_id = 2 + + def __call__(self, prompt, add_special_tokens=False): + return SimpleNamespace(input_ids=prompt) + + +def test_tokenizer_sets_image_block_span(): + tokenizer = Gemma4Tokenizer( + _Tokenizer(), + { + "image_token_id": 90, + "boi_token_id": 91, + "eoi_token_id": 92, + "vision_soft_tokens_per_image": 3, + }, + ) + image = ImageItem(type="image_size", data=(1, 1)) + image.token_id = 100 + image.token_num = 3 + + result = tokenizer.encode( + [7, 90, 90, 92, 8], + SimpleNamespace(images=[image]), + ) + + assert result == [7, 91, 100, 101, 102, 92, 8] + assert image.start_idx == 2 + assert image.block_start_idx == 2 + assert image.block_end_idx == 5 + assert image.to_dict()["block_start_idx"] == 2 + assert image.to_dict()["block_end_idx"] == 5 From c8f53bebe1bb1bea92174b4bda6e3f12c34612de Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 3 Sep 2026 06:00:10 +0000 Subject: [PATCH 154/214] support fp8xfp8 mega moe --- .../fused_moe/grouped_fused_moe_ep.py | 125 +++++++++++------- .../quantization/fp8act_quant_kernel.py | 42 ++++++ lightllm/common/quantization/deepgemm.py | 8 ++ lightllm/distributed/communication_op.py | 106 ++++++++++----- .../layer_infer/transformer_layer_infer.py | 6 +- .../layer_infer/transformer_layer_infer.py | 6 +- .../layer_infer/transformer_layer_infer.py | 6 +- .../mode_backend/ep_balance_monitor.py | 6 +- lightllm/utils/device_utils.py | 5 + 9 files changed, 219 insertions(+), 91 deletions(-) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 1c6feea31b..a8b28ceaa3 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -11,6 +11,7 @@ silu_and_mul_masked_post_quant_fwd, ) from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import ( + lightllm_per_token_group_quant_fp8, per_token_group_quant_fp8, ) from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_expanded_layout_kernels import ( @@ -30,7 +31,6 @@ from lightllm.utils.tensor_buffer_manager import TensorBufferManager logger = init_logger(__name__) -_MEGA_MOE_STATES: Dict[Tuple[int, int, int, int], Dict[str, Any]] = {} SUPPORTED_EP_EXPERT_DTYPES = ("fp8w8a8-b128-deepgemm", "fp4fp8-b32-deepgemm") @@ -48,8 +48,8 @@ def get_ep_num_sms() -> int: return getattr(dist_group_manager, "ep_num_sms", None) or 0 -def use_sm100_mega_moe(quant_method: Any) -> bool: - return is_sm100_gpu() and quant_method.method_name == "fp4fp8-b32-deepgemm" +def use_mega_moe(quant_method: Any) -> bool: + return getattr(quant_method, "mega_moe_mma_type", None) is not None def check_ep_expert_dtype(quant_method: Any): @@ -98,23 +98,28 @@ def masked_group_gemm( return gemm_out_b -def _get_mega_moe_cache_state(w13: Any, w2: Any): - state_key = ( - w13.weight.data_ptr(), - w13.weight_scale.data_ptr(), - w2.weight.data_ptr(), - w2.weight_scale.data_ptr(), +def _get_mega_moe_weights(w13: Any, w2: Any, mma_type: str): + weights = getattr(w13, "_mega_moe_weights", None) + if weights is not None: + return weights + transform_kwargs = {"mma_type": mma_type} if mma_type == "fp8xfp8" else {} + weights = deep_gemm.transform_weights_for_mega_moe( + (w13.weight, w13.weight_scale), + (w2.weight, w2.weight_scale), + **transform_kwargs, ) - return _MEGA_MOE_STATES.setdefault(state_key, {}) - - -def _get_mega_moe_weights(w13: Any, w2: Any, state: Dict[str, Any]): - if "weight_cache" not in state: - state["weight_cache"] = deep_gemm.transform_weights_for_mega_moe( - (w13.weight, w13.weight_scale), - (w2.weight, w2.weight_scale), - ) - return state["weight_cache"] + if mma_type == "fp8xfp8": + # Keep the transformed layout in the preallocated weight storage so we do not retain a second + # full copy of the expert weights. Skip copy_ when DeepGEMM already returned an alias. + for target, transformed in zip( + (w13.weight, w13.weight_scale, w2.weight, w2.weight_scale), + (*weights[0], *weights[1]), + ): + if target.data_ptr() != transformed.data_ptr(): + target.copy_(transformed) + weights = ((w13.weight, w13.weight_scale), (w2.weight, w2.weight_scale)) + w13._mega_moe_weights = weights + return weights def _get_mega_moe_cumulative_stats(num_local_experts: int, device: torch.device, state: Dict[str, Any]): @@ -125,6 +130,16 @@ def _get_mega_moe_cumulative_stats(num_local_experts: int, device: torch.device, return stats +def prepare_mega_moe_weights(w13: Any, w2: Any, quant_method: Any): + mma_type = dist_group_manager.ep_mega_moe_mma_type + if dist_group_manager.ep_mega_moe_quant_method != quant_method.method_name: + quant_method.mega_moe_mma_type = None + return + quant_method.mega_moe_mma_type = mma_type + if mma_type == "fp8xfp8": + _get_mega_moe_weights(w13, w2, mma_type) + + def mega_moe_impl( hidden_states: torch.Tensor, w13: Any, @@ -132,17 +147,17 @@ def mega_moe_impl( topk_weights: torch.Tensor, topk_ids: torch.Tensor, quant_method: Any, + mma_type: str, clamp_limit: Optional[float] = None, alloc_tensor_func: Callable = torch.empty, ): - if not (HAS_DEEPGEMM and hasattr(deep_gemm, "fp8_fp4_mega_moe")): - raise RuntimeError("deep_gemm does not provide fp8-fp4 Mega MoE kernel") - - from deep_gemm.utils import per_token_cast_to_fp8 + kernel_name = "fp8_fp8_mega_moe" if mma_type == "fp8xfp8" else "fp8_fp4_mega_moe" + if not (HAS_DEEPGEMM and hasattr(deep_gemm, kernel_name)): + raise RuntimeError(f"deep_gemm does not provide {kernel_name} Mega MoE kernel") buffer = getattr(dist_group_manager, "ep_mega_moe_buffer", None) if buffer is None: - raise RuntimeError("SM100 Mega MoE requires dist_group_manager.ep_mega_moe_buffer to be initialized") + raise RuntimeError("Mega MoE requires dist_group_manager.ep_mega_moe_buffer to be initialized") num_tokens = hidden_states.shape[0] if num_tokens > buffer.num_max_tokens_per_rank: @@ -150,22 +165,42 @@ def mega_moe_impl( f"Mega MoE got {num_tokens} tokens, exceeding num_max_tokens_per_rank={buffer.num_max_tokens_per_rank}" ) - qinput_tensor = per_token_cast_to_fp8( - hidden_states, - use_ue8m0=True, - gran_k=quant_method.block_size, - use_packed_ue8m0=True, - ) - state = _get_mega_moe_cache_state(w13, w2) - l1_weights, l2_weights = _get_mega_moe_weights(w13, w2, state) - stats = _get_mega_moe_cumulative_stats(w13.weight.shape[0], hidden_states.device, state) - buffer.x[:num_tokens].copy_(qinput_tensor[0]) - buffer.x_sf[:num_tokens].copy_(qinput_tensor[1]) - buffer.topk_idx[:num_tokens].copy_(topk_ids) - buffer.topk_weights[:num_tokens].copy_(topk_weights) + if mma_type == "fp8xfp8": + lightllm_per_token_group_quant_fp8( + x=hidden_states, + group_size=quant_method.block_size, + x_q=buffer.x[:num_tokens], + x_s=buffer.x_sf[:num_tokens], + eps=1e-4, + dtype=buffer.x.dtype, + topk_ids=topk_ids, + topk_weights=topk_weights, + topk_ids_out=buffer.topk_idx[:num_tokens], + topk_weights_out=buffer.topk_weights[:num_tokens], + ) + else: + from deep_gemm.utils import per_token_cast_to_fp8 + qinput_tensor = per_token_cast_to_fp8( + hidden_states, + use_ue8m0=True, + gran_k=quant_method.block_size, + use_packed_ue8m0=True, + ) + buffer.x[:num_tokens].copy_(qinput_tensor[0]) + buffer.x_sf[:num_tokens].copy_(qinput_tensor[1]) + buffer.topk_idx[:num_tokens].copy_(topk_ids) + buffer.topk_weights[:num_tokens].copy_(topk_weights) + + l1_weights, l2_weights = _get_mega_moe_weights(w13, w2, mma_type) + state = getattr(w13, "_mega_moe_state", None) + if state is None: + state = {} + w13._mega_moe_state = state + stats = _get_mega_moe_cumulative_stats(w13.weight.shape[0], hidden_states.device, state) output = alloc_tensor_func(hidden_states.shape, device=hidden_states.device, dtype=hidden_states.dtype) - deep_gemm.fp8_fp4_mega_moe( + kernel = getattr(deep_gemm, kernel_name) + kernel( output, l1_weights, l2_weights, @@ -182,16 +217,6 @@ def quantize_fused_experts_input( quant_method: Any, ): check_ep_expert_dtype(quant_method) - if use_sm100_mega_moe(quant_method): - from deep_gemm.utils import per_token_cast_to_fp8 - - return per_token_cast_to_fp8( - hidden_states, - use_ue8m0=True, - gran_k=quant_method.block_size, - use_packed_ue8m0=True, - ) - block_size_k = 0 if w13.weight.ndim == 3: block_size_k = w13.weight.shape[2] // w13.weight_scale.shape[2] @@ -214,7 +239,8 @@ def fused_experts( ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): check_ep_expert_dtype(quant_method) - if use_sm100_mega_moe(quant_method): + mma_type = getattr(quant_method, "mega_moe_mma_type", None) + if mma_type is not None: return mega_moe_impl( hidden_states, w13, @@ -222,6 +248,7 @@ def fused_experts( topk_weights, topk_idx, quant_method, + mma_type, clamp_limit=clamp_limit, alloc_tensor_func=alloc_tensor_func, ) diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py index a439a01f9d..c3a5b9f634 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py @@ -36,9 +36,18 @@ def _per_token_group_quant_fp8( xs_n, xs_stride_m, xs_stride_n, + topk_ids_ptr, + topk_weights_ptr, + topk_ids_out_ptr, + topk_weights_out_ptr, + num_topk, + topk_row_stride, + topk_out_row_stride, BLOCK: tl.constexpr, + TOPK_BLOCK: tl.constexpr, NEED_MASK: tl.constexpr, USE_UE8M0_SCALE: tl.constexpr, + COPY_TOPK: tl.constexpr, ): g_id = tl.program_id(0) y_ptr += g_id * y_stride @@ -68,6 +77,16 @@ def _per_token_group_quant_fp8( tl.store(y_q_ptr + cols, y_q, mask=mask) tl.store(y_s_ptr, y_s) + if COPY_TOPK: + topk_cols = tl.arange(0, TOPK_BLOCK) + topk_mask = (col_id == 0) & (topk_cols < num_topk) + topk_offsets = row_id * topk_row_stride + topk_cols + topk_out_offsets = row_id * topk_out_row_stride + topk_cols + topk_ids = tl.load(topk_ids_ptr + topk_offsets, mask=topk_mask) + topk_weights = tl.load(topk_weights_ptr + topk_offsets, mask=topk_mask) + tl.store(topk_ids_out_ptr + topk_out_offsets, topk_ids, mask=topk_mask) + tl.store(topk_weights_out_ptr + topk_out_offsets, topk_weights, mask=topk_mask) + def lightllm_per_token_group_quant_fp8( x: torch.Tensor, @@ -77,6 +96,10 @@ def lightllm_per_token_group_quant_fp8( eps: float = 1e-10, dtype: torch.dtype = torch.float8_e4m3fn, use_ue8m0_scales: bool = False, + topk_ids: Optional[torch.Tensor] = None, + topk_weights: Optional[torch.Tensor] = None, + topk_ids_out: Optional[torch.Tensor] = None, + topk_weights_out: Optional[torch.Tensor] = None, ): """group-wise, per-token quantization on input tensor `x`. Args: @@ -103,6 +126,16 @@ def lightllm_per_token_group_quant_fp8( # heuristics for number of warps num_warps = min(max(BLOCK // 256, 1), 8) num_stages = 1 + copy_topk = topk_ids is not None + if copy_topk: + num_topk = topk_ids.shape[-1] + topk_block = triton.next_power_of_2(num_topk) + topk_row_stride = topk_ids.stride(0) + topk_out_row_stride = topk_ids_out.stride(0) + else: + topk_ids = topk_weights = topk_ids_out = topk_weights_out = x + topk_block = 1 + num_topk = topk_row_stride = topk_out_row_stride = 0 _per_token_group_quant_fp8[(M,)]( x, x_q, @@ -115,9 +148,18 @@ def lightllm_per_token_group_quant_fp8( xs_n=xs_n, xs_stride_m=xs_stride_m, xs_stride_n=xs_stride_n, + topk_ids_ptr=topk_ids, + topk_weights_ptr=topk_weights, + topk_ids_out_ptr=topk_ids_out, + topk_weights_out_ptr=topk_weights_out, + num_topk=num_topk, + topk_row_stride=topk_row_stride, + topk_out_row_stride=topk_out_row_stride, BLOCK=BLOCK, + TOPK_BLOCK=topk_block, NEED_MASK=BLOCK != group_size, USE_UE8M0_SCALE=use_ue8m0_scales, + COPY_TOPK=copy_topk, num_warps=num_warps, num_stages=num_stages, ) diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 1cb134a7e1..436d0bc50f 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -24,10 +24,18 @@ def __init__(self): self.cache_manager = g_cache_manager assert HAS_DEEPGEMM, "deepgemm is not installed, you can't use quant api of it" + self.mega_moe_mma_type = None def quantize(self, weight: torch.Tensor, output: WeightPack): raise NotImplementedError("Not implemented") + def finalize_moe_weight(self, moe_weight): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( + prepare_mega_moe_weights, + ) + + prepare_mega_moe_weights(moe_weight.w13, moe_weight.w2, self) + def apply( self, input_tensor: torch.Tensor, diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 3fedac78d6..08015789a6 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -24,8 +24,8 @@ from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union from lightllm.utils.log_utils import init_logger -from lightllm.utils.device_utils import has_nvlink from lightllm.utils.envs_utils import ( + enable_env_vars, get_env_start_args, get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -37,10 +37,17 @@ create_new_group_for_current_dp, create_dp_special_inter_group, ) -from lightllm.utils.device_utils import get_device_sm_count, is_sm100_gpu +from lightllm.utils.device_utils import ( + get_device_sm_count, + has_nvlink, + is_sm90_gpu, + is_sm100_gpu, +) from lightllm.utils.torch_dtype_utils import get_torch_dtype logger = init_logger(__name__) +FP8_MOE_QUANT_METHOD = "fp8w8a8-b128-deepgemm" +FP4_MOE_QUANT_METHOD = "fp4fp8-b32-deepgemm" def get_deep_ep_prefill_moe_workspace_size( @@ -163,6 +170,8 @@ def __init__(self): self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None + self.ep_mega_moe_mma_type = None + self.ep_mega_moe_quant_method = None self.ep_num_sms = None def __len__(self): @@ -230,9 +239,9 @@ def new_deepep_group( """初始化 DeepEP 通信组以及当前模型实际需要的 MoE buffer。 ``expert_quant_method_names`` 是各 MoE 层最终绑定的 quant method 名称集合。 - 同一个模型可能逐层混用 FP4 和 FP8:SM100 FP4 层走 Mega MoE,其他层走 - DeepEP legacy low-latency 路径。这里只为实际存在的执行路径分配 buffer, - 避免为未使用的路径长期占用显存。 + 同一个模型可能逐层混用多种 expert quant method:满足约束的 SM100 FP4 和 + SM90 FP8 层走 Mega MoE,其他层走 DeepEP legacy 路径。这里只为实际存在的 + 执行路径分配 buffer,避免为未使用的路径长期占用显存。 """ args = get_env_start_args() enable_ep_moe = args.enable_ep_moe @@ -245,6 +254,8 @@ def new_deepep_group( self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None + self.ep_mega_moe_mma_type = None + self.ep_mega_moe_quant_method = None self.ep_num_sms = None return assert HAS_DEEPEP, "deep_ep is required for expert parallelism" @@ -261,7 +272,8 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - self.ll_num_experts = n_routed_experts + get_redundancy_expert_num() * global_world_size + redundancy_expert_num = get_redundancy_expert_num() + self.ll_num_experts = n_routed_experts + redundancy_expert_num * global_world_size self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, @@ -277,25 +289,45 @@ def new_deepep_group( if not expert_quant_method_names: raise ValueError("No valid MoE quant method was found while initializing DeepEP buffers") - mega_moe_quant_method = "fp4fp8-b32-deepgemm" - is_sm100 = is_sm100_gpu() - - # Buffer 选择规则: - # 1. 非 SM100 不支持 Mega MoE,只初始化 legacy low-latency buffer; - # 2. SM100 全部 MoE 层为 FP4,只初始化 Mega MoE buffer; - # 3. SM100 全部 MoE 层为 FP8,只初始化 legacy low-latency buffer; - # 4. SM100 逐层混合 FP4/FP8,两套 buffer 都要初始化。 - if is_sm100: - # 只要存在一个 FP4 MoE 层,就需要 Mega MoE buffer;只要存在一个非 FP4 - # MoE 层,就需要 legacy low-latency buffer。FP4/FP8 逐层混用时两者都会初始化。 - has_mega_moe_layer = mega_moe_quant_method in expert_quant_method_names - has_legacy_moe_layer = any( - method_name != mega_moe_quant_method for method_name in expert_quant_method_names - ) - enable_mega_moe_buffer = has_mega_moe_layer - else: - enable_mega_moe_buffer = False - has_legacy_moe_layer = True + self.ep_mega_moe_mma_type = None + self.ep_mega_moe_quant_method = None + if is_sm100_gpu() and FP4_MOE_QUANT_METHOD in expert_quant_method_names: + self.ep_mega_moe_mma_type = "fp8xfp4" + self.ep_mega_moe_quant_method = FP4_MOE_QUANT_METHOD + elif ( + enable_env_vars("LIGHTLLM_ENABLE_SM90_FP8_MEGA_MOE") + and is_sm90_gpu() + and FP8_MOE_QUANT_METHOD in expert_quant_method_names + and redundancy_expert_num == 0 + ): + self.ep_mega_moe_mma_type = "fp8xfp8" + self.ep_mega_moe_quant_method = FP8_MOE_QUANT_METHOD + if self.ep_mega_moe_mma_type == "fp8xfp8": + import deep_gemm + + fallback_reason = None + if not hasattr(deep_gemm, "fp8_fp8_mega_moe") or not hasattr( + getattr(deep_gemm, "_C", None), "fp8_fp8_mega_moe" + ): + fallback_reason = ( + "the loaded DeepGEMM Python package and extension do not both provide fp8_fp8_mega_moe " + f"({getattr(deep_gemm, '__file__', '')})" + ) + elif getattr(args, "nnodes", 1) != 1: + fallback_reason = "Mega MoE only supports a single-node expert-parallel group" + elif getattr(args, "enable_rl", False): + fallback_reason = "online expert-weight updates require the canonical non-interleaved layout" + elif not has_nvlink(): + fallback_reason = "NVLink is unavailable" + if fallback_reason is not None: + logger.warning("Disable SM90 FP8 Mega MoE and use legacy DeepEP because %s", fallback_reason) + self.ep_mega_moe_mma_type = None + self.ep_mega_moe_quant_method = None + + enable_mega_moe_buffer = self.ep_mega_moe_mma_type is not None + has_legacy_moe_layer = not enable_mega_moe_buffer or any( + method_name != self.ep_mega_moe_quant_method for method_name in expert_quant_method_names + ) enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" @@ -345,14 +377,19 @@ def new_deepep_group( device=torch.device("cuda", torch.cuda.current_device()), ) + theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_num_experts, num_experts_per_tok) + deepep_sms = 0 if self.ep_mega_moe_mma_type == "fp8xfp8" and not has_legacy_moe_layer else theoretical_sms + self._set_num_sms_for_deep_gemm(deepep_sms) + if enable_mega_moe_buffer: - # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy - # low-latency buffer,因此纯 FP4 模型无需承担后者的大块 RDMA 显存。 if moe_intermediate_size is None: - raise ValueError("SM100 Mega MoE requires moe_intermediate_size or intermediate_size in model config") + raise ValueError("Mega MoE requires moe_intermediate_size or intermediate_size in model config") import deep_gemm + mega_buffer_kwargs = ( + {"mma_type": self.ep_mega_moe_mma_type} if self.ep_mega_moe_mma_type == "fp8xfp8" else {} + ) self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, self.ll_num_experts, @@ -360,17 +397,17 @@ def new_deepep_group( num_experts_per_tok, self.ll_hidden, moe_intermediate_size, + **mega_buffer_kwargs, ) logger.info( "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " - "mega_moe=%s, expert_quant_method_names=%s", + "mega_moe=%s, mega_moe_mma_type=%s, expert_quant_method_names=%s", enable_low_latency_buffer, self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, + self.ep_mega_moe_mma_type, sorted(expert_quant_method_names), ) - theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_num_experts, num_experts_per_tok) - self._set_num_sms_for_deep_gemm(theoretical_sms) def _set_num_sms_for_deep_gemm(self, deepep_sms: int): try: @@ -384,8 +421,13 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): self.ep_num_sms = deepep_sms if self.ep_low_latency_buffer is not None: deep_ep.Buffer.set_num_sms(deepep_sms - deepep_sms % 2) - set_num_sms(max(device_sms - deepep_sms, 2)) + deep_gemm_sms = max(device_sms - deepep_sms, 2) + if self.ep_mega_moe_mma_type == "fp8xfp8": + deep_gemm_sms -= deep_gemm_sms % 2 + set_num_sms(deep_gemm_sms) except BaseException as e: + if self.ep_mega_moe_mma_type is not None: + raise RuntimeError("Failed to reserve a fixed SM pool before allocating the Mega MoE buffer") from e logger.warning(f"set num sms for deep_gemm failed: {e}") def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index 3254031056..b2cbb737bf 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -8,7 +8,7 @@ from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.deepseek2.infer_struct import Deepseek2InferStateInfo from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - use_sm100_mega_moe, + use_mega_moe, ) from functools import partial from lightllm.models.llama.yarn_rotary_utils import get_deepseek_mscale @@ -300,7 +300,7 @@ def overlap_tpsp_token_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): + if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -426,7 +426,7 @@ def overlap_tpsp_context_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): + if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 166cc0efb0..eb96389801 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -2,7 +2,7 @@ import triton from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_sm100_mega_moe +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight @@ -151,7 +151,7 @@ def overlap_tpsp_context_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): + if not self.enable_ep_moe or use_mega_moe(experts.quant_method): input_embdings = self.context_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.context_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 @@ -256,7 +256,7 @@ def overlap_tpsp_token_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): + if not self.enable_ep_moe or use_mega_moe(experts.quant_method): input_embdings = self.token_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.token_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 diff --git a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py index 7311c4d141..aa45940440 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -7,7 +7,7 @@ from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - use_sm100_mega_moe, + use_mega_moe, ) from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.envs_utils import get_env_start_args @@ -138,7 +138,7 @@ def overlap_tpsp_token_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): + if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -250,7 +250,7 @@ def overlap_tpsp_context_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): + if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py index 107f4ba504..cf1074730a 100644 --- a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py +++ b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py @@ -31,7 +31,11 @@ def should_enable_ep_balance_monitor(args) -> bool: - if args.enable_prefill_cudagraph or is_sm100_gpu(): + if ( + args.enable_prefill_cudagraph + or is_sm100_gpu() + or getattr(dist_group_manager, "ep_mega_moe_buffer", None) is not None + ): return False return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" diff --git a/lightllm/utils/device_utils.py b/lightllm/utils/device_utils.py index 58bff90560..7135e3f232 100644 --- a/lightllm/utils/device_utils.py +++ b/lightllm/utils/device_utils.py @@ -45,6 +45,11 @@ def is_sm100_gpu(): return torch.cuda.get_device_capability()[0] == 10 +@lru_cache(maxsize=None) +def is_sm90_gpu(): + return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 9 + + @lru_cache(maxsize=None) def get_device_sm_regs_num(): import triton From f4f4d68d0cee945a0892393518a3720c89ef69d2 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Fri, 4 Sep 2026 02:54:33 +0000 Subject: [PATCH 155/214] fix: account for DSpark SWA scratch pages --- .../server/router/model_infer/infer_batch.py | 4 ++++ test/unit/test_deepseek_v4_dspark.py | 18 ++++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index ee99cf2720..8a980ccee7 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1123,6 +1123,10 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: cur_seq_len = seq_len + step if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: swa_page_num += 1 + + # DSpark proposal allocates one private SWA scratch page per request. + if self.args.mtp_mode == "dspark": + swa_page_num += 1 return swa_page_num, c4_page_num, c128_slot_num diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index e4432852b2..fc260d93bc 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -8,6 +8,24 @@ from lightllm.utils.config_utils import get_deepseek_v4_compress_rates +def test_dspark_decode_admission_reserves_one_swa_scratch_page(): + from lightllm.server.router.model_infer.infer_batch import InferReq + + req = InferReq.__new__(InferReq) + req.args = SimpleNamespace(mtp_mode="eagle") + req.mtp_step = 1 + req.dsv4_swa_page_size = 128 + req.dsv4_c4_page_size = 64 + req.dsv4_has_c128 = True + req.get_cur_total_len = lambda: 10 + + normal_need = req.get_dsv4_decode_need_page_and_slot_num() + req.args.mtp_mode = "dspark" + dspark_need = req.get_dsv4_decode_need_page_and_slot_num() + + assert dspark_need == (normal_need[0] + 1, normal_need[1], normal_need[2]) + + @pytest.mark.parametrize("mtp_step", [1, 4, 5]) def test_deepseek_v4_dspark_runtime_width_follows_mtp_step(monkeypatch, mtp_step): from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel From a7f9c3686a9da7de500e69ae2c84a42e165f7251 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 4 Sep 2026 04:18:25 +0000 Subject: [PATCH 156/214] fix: avoid EAGLE SWA page overcount in DSpark admission --- .../server/router/model_infer/infer_batch.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 8a980ccee7..9ba8163b7b 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -1117,16 +1117,16 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: if self.dsv4_has_c128 and cur_seq_len % 128 == 0: c128_slot_num += 1 - # EAGLE draft forwards after the first one consume newly appended draft-only rows. - # The DeepSeek-V4 MTP draft layer is compress_ratio=0, so these rows need only SWA. - for step in range(self.mtp_step + 1, self.mtp_step * 2): - cur_seq_len = seq_len + step - if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: - swa_page_num += 1 - - # DSpark proposal allocates one private SWA scratch page per request. if self.args.mtp_mode == "dspark": + # DSpark proposal allocates one private SWA scratch page per request. swa_page_num += 1 + else: + # EAGLE draft forwards after the first one consume newly appended draft-only rows. + # The DeepSeek-V4 MTP draft layer is compress_ratio=0, so these rows need only SWA. + for step in range(self.mtp_step + 1, self.mtp_step * 2): + cur_seq_len = seq_len + step + if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: + swa_page_num += 1 return swa_page_num, c4_page_num, c128_slot_num From 7229e567526c9428070e03bd025c138137d4390a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 4 Sep 2026 04:34:39 +0000 Subject: [PATCH 157/214] clean disk cache on shutdown --- lightllm/server/api_start.py | 9 +- .../multi_level_kv_cache/disk_cache_worker.py | 10 +- .../server/multi_level_kv_cache/manager.py | 6 +- lightllm/utils/start_utils.py | 111 +++++++++++------- 4 files changed, 85 insertions(+), 51 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 1328f02ef2..6bce2ceb21 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -1,5 +1,6 @@ import multiprocessing as mp import os +import tempfile import uuid import subprocess import math @@ -412,6 +413,12 @@ def _launch_subprocesses(args: StartArgs): ], ) + instance_disk_cache_dir = None + if args.enable_cpu_cache and args.enable_disk_cache: + cache_base_dir = args.disk_cache_dir or tempfile.gettempdir() + instance_disk_cache_dir = os.path.join(cache_base_dir, f"lightllm_disk_cache_{get_unique_server_name()}") + process_manager.register_disk_cache_dir(instance_disk_cache_dir) + if args.enable_cpu_cache: from .multi_level_kv_cache.manager import start_multi_level_kv_cache_manager @@ -419,7 +426,7 @@ def _launch_subprocesses(args: StartArgs): start_funcs=[ start_multi_level_kv_cache_manager, ], - start_args=[(args,)], + start_args=[(args, instance_disk_cache_dir)], ) process_manager.start_submodule_processes( diff --git a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py index 07a0be319b..aff400753d 100644 --- a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py +++ b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py @@ -1,12 +1,10 @@ import os -import tempfile import time import math from dataclasses import dataclass -from typing import List, Optional +from typing import List import torch -from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from .cpu_cache_client import CpuKvCacheClient @@ -37,7 +35,7 @@ def __init__( self, disk_cache_storage_size: float, cpu_cache_client: CpuKvCacheClient, - disk_cache_dir: Optional[str] = None, + cache_dir: str, ): self.cpu_cache_client = cpu_cache_client self._pages_all_idle = False @@ -50,10 +48,6 @@ def __init__( # 读写同时进行时,分配8线程用来写,16线程用来读 max_concurrent_write_tasks = 8 - if disk_cache_dir: - cache_dir = os.path.join(disk_cache_dir, f"lightllm_disk_cache_{get_unique_server_name()}") - else: - cache_dir = os.path.join(tempfile.gettempdir(), f"lightllm_disk_cache_{get_unique_server_name()}") os.makedirs(cache_dir, exist_ok=True) cache_file = os.path.join(cache_dir, "cache_file") diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index 4f8a6d007c..f2b4a576a7 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -27,6 +27,7 @@ class MultiLevelKVCacheManager: def __init__( self, args: StartArgs, + instance_disk_cache_dir, ): self.args: StartArgs = args ports = get_shm_port_args() @@ -59,7 +60,7 @@ def __init__( self.disk_cache_worker = DiskCacheWorker( disk_cache_storage_size=self.args.disk_cache_storage_size, cpu_cache_client=self.cpu_cache_client, - disk_cache_dir=self.args.disk_cache_dir, + cache_dir=instance_disk_cache_dir, ) self.disk_cache_thread = threading.Thread(target=self.disk_cache_worker.run, daemon=True) self.disk_cache_thread.start() @@ -262,7 +263,7 @@ def recv_loop(self): return -def start_multi_level_kv_cache_manager(args, pipe_writer): +def start_multi_level_kv_cache_manager(args, instance_disk_cache_dir, pipe_writer): # 注册graceful 退出的处理 graceful_registry(inspect.currentframe().f_code.co_name) setproctitle.setproctitle(f"lightllm::{get_unique_server_name()}::multi_level_kv_cache") @@ -271,6 +272,7 @@ def start_multi_level_kv_cache_manager(args, pipe_writer): try: manager = MultiLevelKVCacheManager( args=args, + instance_disk_cache_dir=instance_disk_cache_dir, ) except Exception as e: logger.exception(f"start multi_level_kv_cache_manager has exception {str(e)}") diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index 65b6f5f41b..f66b5a2464 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -1,4 +1,5 @@ import os +import shutil import signal import subprocess import sys @@ -15,38 +16,56 @@ class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} + self.disk_cache_dir = None + + def register_disk_cache_dir(self, cache_dir): + self.disk_cache_dir = cache_dir def start_submodule_processes(self, start_funcs=[], start_args=[]): assert len(start_funcs) == len(start_args) pipe_readers = [] processes = [] - for start_func, start_arg in zip(start_funcs, start_args): - pipe_reader, pipe_writer = mp.Pipe(duplex=False) - process = mp.Process( - target=start_func, - args=start_arg + (pipe_writer,), - ) - process.start() - pipe_readers.append(pipe_reader) - processes.append(process) - - # Wait for all processes to initialize - for index, pipe_reader in enumerate(pipe_readers): - init_state = pipe_reader.recv() - if init_state != "init ok": - logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") - for proc in processes: - proc.kill() - sys.exit(1) - else: + try: + for start_func, start_arg in zip(start_funcs, start_args): + pipe_reader, pipe_writer = mp.Pipe(duplex=False) + process = mp.Process( + target=start_func, + args=start_arg + (pipe_writer,), + ) + try: + process.start() + finally: + pipe_writer.close() + pipe_readers.append(pipe_reader) + processes.append(process) + + # Wait for all processes to initialize + for index, pipe_reader in enumerate(pipe_readers): + try: + init_state = pipe_reader.recv() + finally: + pipe_reader.close() + if init_state != "init ok": + logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") + raise SystemExit(1) logger.info(f"init func {start_funcs[index].__name__} : {str(init_state)}") - assert all([proc.is_alive() for proc in processes]) - processes = [psutil.Process(proc.pid) for proc in processes] - self.processes.extend(processes) - self.process_names.update((process, process.name()) for process in processes) - return processes + assert all([proc.is_alive() for proc in processes]) + managed_processes = [psutil.Process(proc.pid) for proc in processes] + managed_process_names = {process: process.name() for process in managed_processes} + except BaseException: + for proc in processes: + if proc.is_alive(): + proc.kill() + for proc in processes: + proc.join() + self.terminate_all_processes() + raise + + self.processes.extend(managed_processes) + self.process_names.update(managed_process_names) + return managed_processes def register_process_tree(self, root_process): """Add persistent LightLLM descendants to supervision. @@ -91,9 +110,19 @@ def kill_recursive(proc): kill_recursive(proc) proc.wait() + # LightMem owns files under this directory, so remove it only after the cache process has exited. + if self.disk_cache_dir is not None: + try: + shutil.rmtree(self.disk_cache_dir) + except FileNotFoundError: + pass + except Exception as e: + logger.warning(f"Failed to remove disk cache directory {self.disk_cache_dir}: {e}") + else: + logger.info(f"Removed disk cache directory {self.disk_cache_dir}") + # recover the gpu compute mode - is_enable_mps = get_env_start_args().enable_mps - if is_enable_mps: + if get_env_start_args().enable_mps: from lightllm.utils.device_utils import stop_mps stop_mps() @@ -103,10 +132,11 @@ def setup_signal_handlers(self, http_server_process=None): def signal_handler(sig, _frame): if sig == signal.SIGINT: logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") - if http_server_process is not None: - kill_recursive(http_server_process) - - self.terminate_all_processes() + try: + if http_server_process is not None: + kill_recursive(http_server_process) + finally: + self.terminate_all_processes() logger.info("All processes have been forcefully terminated.") sys.exit(0) @@ -115,16 +145,17 @@ def signal_handler(sig, _frame): else: logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") - if http_server_process is not None and http_server_process.poll() is None: - http_server_process.send_signal(signal.SIGTERM) - try: - http_server_process.wait(timeout=60) - logger.info("HTTP server exited gracefully") - except subprocess.TimeoutExpired: - logger.warning("HTTP server did not exit in time, killing it...") - kill_recursive(http_server_process) - - self.terminate_all_processes() + try: + if http_server_process is not None and http_server_process.poll() is None: + http_server_process.send_signal(signal.SIGTERM) + try: + http_server_process.wait(timeout=60) + logger.info("HTTP server exited gracefully") + except subprocess.TimeoutExpired: + logger.warning("HTTP server did not exit in time, killing it...") + kill_recursive(http_server_process) + finally: + self.terminate_all_processes() logger.info("All processes have been terminated gracefully.") sys.exit(0) From ca4f075860e00649ea73554780b33a417ead21bc Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sat, 5 Sep 2026 01:34:22 +0000 Subject: [PATCH 158/214] synchronize PDL top-k and fail fast on model thread errors --- .../deepseek_v4/triton_kernel/csrc/topk_transform.cu | 7 +++++-- lightllm/server/router/model_infer/model_rpc.py | 3 ++- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu index 53d430243d..e29b57fc95 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu @@ -237,6 +237,11 @@ struct TopKParams { template __global__ void topk_transform_kernel(const __grid_constant__ TopKParams params) { const uint32_t work_id = blockIdx.x; + + // PDL may start this grid before its producer finishes, so wait before the + // first global load from seq_lens. + pdl_wait_primary(); + const uint32_t seq_len = params.seq_lens[work_id]; const auto score_ptr = params.scores + work_id * params.score_stride; const auto page_ptr = params.page_table + work_id * params.page_table_stride; @@ -244,8 +249,6 @@ __global__ void topk_transform_kernel(const __grid_constant__ TopKParams params) const auto raw_indices_ptr = params.raw_indices != nullptr ? params.raw_indices + work_id * kTopK : nullptr; const uint32_t page_bits = params.page_bits; - pdl_wait_primary(); - if (seq_len <= kTopK) { naive_transform(page_ptr, indices_ptr, raw_indices_ptr, seq_len, page_bits); } else { diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 3ae4d4cbc2..04c78e74ab 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -34,7 +34,7 @@ from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.utils.log_utils import init_logger from lightllm.utils.graceful_utils import graceful_registry -from lightllm.utils.process_check import start_parent_check_thread +from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.torch_memory_saver_utils import MemoryTag from lightllm.server.io_struct import RlOpReq, RlOpRsp @@ -180,6 +180,7 @@ def _init_env( socket_path, success_event, ): + install_fatal_thread_excepthook() import lightllm.utils.rpyc_fix_utils as _ # 注册graceful 退出的处理 From 9fb708ff8ce2ede891df531840721548f1d6e2bc Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 6 Sep 2026 01:40:45 +0000 Subject: [PATCH 159/214] fix(dsv4): invalidate C128 mapping before releasing slots --- lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index da12d29b8f..7e7138b759 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -981,8 +981,10 @@ def _evict_compress(self, full_slots: torch.Tensor, mapping: torch.Tensor, alloc valid_slots = slots[valid] if valid_slots.numel() == 0: return - allocator.free(valid_slots) mapping[full_slots[valid]] = -1 + # The allocator's blocking D2H copy must also finish invalidating the + # old mapping before another stream can reuse the returned slots. + allocator.free(valid_slots) return def alloc_c4_pages(self, need_pages: int) -> torch.Tensor: From 7bc0955b20816d407748dcc927f461661609c155 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 6 Sep 2026 01:42:28 +0000 Subject: [PATCH 160/214] fix(pd): handle deep prompt-cache keys iteratively --- .../pd_selector/cache_aware.py | 3 - .../pd_selector/prompt_cache_tree.py | 79 +++++++++---------- 2 files changed, 37 insertions(+), 45 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index 02319dff5a..175963ac37 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -53,8 +53,6 @@ class CacheAwareConfig: evict_node_batch: int = 10_000 # 每隔 sample_stride 个字符抽 1 个作为前缀树 key,降低匹配开销与内存。 sample_stride: int = 512 - # 初始化前缀树时通过 sys.setrecursionlimit 调大 Python 调用栈深度。 - recursion_limit: int = 4000 class BalanceRelThresholdController: @@ -108,7 +106,6 @@ def __init__(self, config: Optional[CacheAwareConfig] = None) -> None: sample_stride=self.config.sample_stride, max_node_count=self.config.max_node_count, evict_node_batch=self.config.evict_node_batch, - recursion_limit=self.config.recursion_limit, ) self.balance_rel_threshold_controller = BalanceRelThresholdController() diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py b/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py index 859b7a7329..5833ede354 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py @@ -1,6 +1,5 @@ from __future__ import annotations -import sys import time from dataclasses import dataclass, field from threading import Lock, RLock @@ -39,8 +38,7 @@ class PromptCacheTree: """ 用于 cache-aware 选点的 prompt 前缀缓存树。 - prompt 先按 sample_stride 抽稀成 key,再递归按单字符建树/匹配。 - recursion_limit 在初始化时通过 sys.setrecursionlimit 调大 Python 调用栈深度。 + prompt 先按 sample_stride 抽稀成 key,再按单字符建树/匹配。 整棵树节点数有上限;超限时按 LRU 从叶节点批量删除。 """ @@ -49,14 +47,12 @@ def __init__( sample_stride: int = 512, max_node_count: int = 1_000_000, evict_node_batch: int = 10_000, - recursion_limit: int = 4000, ) -> None: """ Args: sample_stride: 每隔多少个字符抽 1 个作为 trie key。 max_node_count: 树中允许的最大节点数(不含 root);超限时触发 LRU 驱逐。 evict_node_batch: 每次驱逐时在超出量基础上额外腾出的节点缓冲数。 - recursion_limit: 初始化时通过 sys.setrecursionlimit 设置的调用栈深度上限。 """ if sample_stride < 1: raise ValueError(f"sample_stride must be >= 1, got {sample_stride}") @@ -64,14 +60,9 @@ def __init__( raise ValueError(f"max_node_count must be >= 0, got {max_node_count}") if evict_node_batch < 1: raise ValueError(f"evict_node_batch must be >= 1, got {evict_node_batch}") - if recursion_limit < 1: - raise ValueError(f"recursion_limit must be >= 1, got {recursion_limit}") self.sample_stride = sample_stride self.max_node_count = max_node_count self.evict_node_batch = evict_node_batch - self.recursion_limit = recursion_limit - if recursion_limit > sys.getrecursionlimit(): - sys.setrecursionlimit(recursion_limit) self.root = _PromptCacheNode() self._node_count = 0 self._leaf_lru: SortedDict[int, _PromptCacheNode] = SortedDict() @@ -105,32 +96,37 @@ def insert(self, text: str, prefill_node: str) -> None: self._evict_if_needed() def _insert_at(self, node: _PromptCacheNode, key: str, depth: int, prefill_node: str) -> None: + path = [] try: - if depth >= len(key): - return - - ch = key[depth] - child = node.children.get(ch) - if child is None: - child = _PromptCacheNode( - parent=node, - edge_char=ch, - last_insert_time=time.monotonic(), - ) - child.last_time_mark = self._gen_time_mark() - node.children[ch] = child - self._node_count += 1 - - self._insert_at(child, key, depth + 1, prefill_node) + while True: + path.append(node) + if depth >= len(key): + break + + ch = key[depth] + child = node.children.get(ch) + if child is None: + child = _PromptCacheNode( + parent=node, + edge_char=ch, + last_insert_time=time.monotonic(), + ) + child.last_time_mark = self._gen_time_mark() + node.children[ch] = child + self._node_count += 1 + + node = child + depth += 1 finally: - if node is not self.root: - node.last_prefill_node = prefill_node - node.last_insert_time = time.monotonic() - if node.last_time_mark in self._leaf_lru: - self._leaf_lru.pop(node.last_time_mark, None) - node.last_time_mark = self._gen_time_mark() - if self._is_leaf(node): - self._leaf_lru[node.last_time_mark] = node + for path_node in reversed(path): + if path_node is not self.root: + path_node.last_prefill_node = prefill_node + path_node.last_insert_time = time.monotonic() + if path_node.last_time_mark in self._leaf_lru: + self._leaf_lru.pop(path_node.last_time_mark, None) + path_node.last_time_mark = self._gen_time_mark() + if self._is_leaf(path_node): + self._leaf_lru[path_node.last_time_mark] = path_node def _gen_time_mark(self) -> int: with self._time_mark_lock: @@ -178,14 +174,13 @@ def prefix_match(self, text: str) -> PromptCacheMatchResult: ) def _match_at(self, node: _PromptCacheNode, key: str, depth: int) -> Tuple[_PromptCacheNode, int]: - if depth >= len(key): - return node, depth - - child = node.children.get(key[depth]) - if child is None: - return node, depth - - return self._match_at(child, key, depth + 1) + while depth < len(key): + child = node.children.get(key[depth]) + if child is None: + break + node = child + depth += 1 + return node, depth def evict_lru_nodes(self) -> int: """节点数超上限时,按 LRU 从叶节点批量删除。 From 02aede0d3b54145b25f526ced89bd1ba89da4a31 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 6 Sep 2026 01:44:25 +0000 Subject: [PATCH 161/214] perf(dsv4): avoid redundant FlashMLA prefill output copy --- .../attention/nsa/dsv4_fp8_flashmla_sparse.py | 25 +++++++++++-------- lightllm/models/deepseek_v4/model.py | 11 ++++---- 2 files changed, 21 insertions(+), 15 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py index 52e38fd4f0..33761d4eff 100644 --- a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py @@ -123,17 +123,22 @@ def prefill_att( dtype=q.dtype, device=q.device, ) - full_out = self.infer_state.dsv4_workspace.flashmla_prefill_full_out[: q.shape[0]] - out.copy_( - self.backend._flashmla_att( - q, - k, - self.infer_state.mem_manager, - nsa_dict, - self._get_sched_meta(nsa_dict["compress_ratio"]), - flashmla_out=full_out, - ) + needs_padding = self.backend.real_q_head_num != self.backend.padded_q_head_num + full_out = ( + self.infer_state.dsv4_workspace.flashmla_prefill_full_out[: q.shape[0]] + if needs_padding + else out.unsqueeze(1) + ) + att_out = self.backend._flashmla_att( + q, + k, + self.infer_state.mem_manager, + nsa_dict, + self._get_sched_meta(nsa_dict["compress_ratio"]), + flashmla_out=full_out, ) + if needs_padding: + out.copy_(att_out) return out diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index a4070dcf05..3e2c47500b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -135,11 +135,12 @@ def _init_att_backend(self): head_dim=self.config["head_dim"], dtype=self.data_type, ) - self.dsv4_workspace.init_flashmla_prefill_full_out( - q_head_num=padded_q_head_num, - head_dim_v=self.config["head_dim"], - dtype=self.data_type, - ) + if padded_q_head_num != real_q_head_num: + self.dsv4_workspace.init_flashmla_prefill_full_out( + q_head_num=padded_q_head_num, + head_dim_v=self.config["head_dim"], + dtype=self.data_type, + ) for layer_infer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): layer_infer.flashmla_q_head_num_ = padded_q_head_num if padded_q_head_num == real_q_head_num: From 8ed3da6e5c36ee4c7a722cce11c1eeb0f1d08f3d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 6 Sep 2026 01:56:26 +0000 Subject: [PATCH 162/214] fix(pd): measure decode KV timeout from request progress --- .../decode_node_impl/decode_trans_process.py | 39 ++++++++++++++++--- 1 file changed, 34 insertions(+), 5 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index ecadca1dac..2a69fdc523 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -141,6 +141,7 @@ def __init__( self.recv_task_group_queue = queue.Queue() self.waiting_dict_lock = threading.Lock() self.waiting_dict: Dict[str, PDChunckedTransTask] = {} + self.request_last_progress_time: Dict[int, float] = {} self.request_page_task_queue = queue.Queue() self.ready_page_task_queue = queue.Queue() self.success_queue = queue.Queue() @@ -214,12 +215,17 @@ def dispatch_task_loop(self): trans_task_group: PDChunckedTransTaskGroup = self.recv_task_group_queue.get() with self.waiting_dict_lock: + has_waiting_task = False for task in trans_task_group.task_list: if task.need_transfer_page(): self.waiting_dict[task.get_key()] = task + has_waiting_task = True else: task.start_trans_time = time.time() self.success_queue.put((None, None, task)) + if has_waiting_task: + request_id = trans_task_group.task_list[0].request_id + self.request_last_progress_time[request_id] = time.time() # up status task = trans_task_group.task_list[0] @@ -243,6 +249,19 @@ def dispatch_task_loop(self): self.up_status_in_queue.put(up_status) + def _pop_waiting_task_for_notify(self, notify_task: PDChunckedTransTask): + with self.waiting_dict_lock: + local_trans_task = self.waiting_dict.pop(notify_task.get_key(), None) + if local_trans_task is None: + return None + + # Decode creates every page task before prefill starts producing pages. + # A matched notify is forward progress for the request, so future pages + # use an idle timeout instead of their original creation time. + self.request_last_progress_time[local_trans_task.request_id] = time.time() + + return local_trans_task + @log_exception def accept_peer_task_loop( self, @@ -291,8 +310,7 @@ def accept_peer_task_loop( # 到了请求页面的阶段 remote_trans_task = notify_obj if remote_trans_task.write_stage == "request": - with self.waiting_dict_lock: - local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) + local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) if local_trans_task is not None: local_trans_task.prefill_agent_name = remote_trans_task.prefill_agent_name local_trans_task.prefill_agent_metadata = remote_trans_task.prefill_agent_metadata @@ -322,8 +340,7 @@ def accept_peer_task_loop( # prefill 写完数据到了 done 阶段 if remote_trans_task.write_stage == "done": - with self.waiting_dict_lock: - local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) + local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) if local_trans_task is not None: local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob @@ -352,9 +369,21 @@ def accept_peer_task_loop( def _check_tasks_time_out(self): with self.waiting_dict_lock: timeout_tasks = [] + pending_request_ids = set() + now = time.time() for key, trans_task in list(self.waiting_dict.items()): - if trans_task.time_out(): + if trans_task.start_trans_time is None: + request_last_progress = self.request_last_progress_time[trans_task.request_id] + is_timeout = now - request_last_progress > trans_task.time_out_secs + else: + is_timeout = trans_task.time_out() + if is_timeout: timeout_tasks.append(self.waiting_dict.pop(key)) + elif trans_task.start_trans_time is None: + pending_request_ids.add(trans_task.request_id) + for request_id in list(self.request_last_progress_time): + if request_id not in pending_request_ids: + self.request_last_progress_time.pop(request_id) for trans_task in timeout_tasks: trans_task.error_info = "time out in accept_peer_task_loop" From aa9b2a34dc9ff310861dc586eae215ebb64dda69 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 6 Sep 2026 12:06:12 +0000 Subject: [PATCH 163/214] fix(api): return 400 for oversized guided JSON schemas --- lightllm/server/core/objs/sampling_params.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 840fa05c35..a136196141 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -176,7 +176,8 @@ class GuidedJsonSchema(ctypes.Structure): def initialize(self, constraint: str, tokenizer): constraint_bytes = constraint.encode("utf-8") - assert len(constraint_bytes) < JSON_SCHEMA_MAX_LENGTH, "Guided json schema is too long." + if len(constraint_bytes) >= JSON_SCHEMA_MAX_LENGTH: + raise ValueError("Guided json schema is too long.") ctypes.memmove(self.constraint, constraint_bytes, len(constraint_bytes)) self.length = len(constraint_bytes) From ef0df4317f6e9db3e19c3366c8e2b2624bc5fb73 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 7 Sep 2026 03:23:13 +0000 Subject: [PATCH 164/214] fix(pd): wait for KV transfer modules to become ready --- .../mode_backend/pd/base_kv_move_manager.py | 6 ++++++ .../decode_kv_move_manager.py | 7 +++++-- .../decode_node_impl/decode_trans_process.py | 5 +++++ .../prefill_kv_move_manager.py | 7 +++++-- .../prefill_trans_process.py | 5 +++++ .../mode_backend/pd/trans_process_obj.py | 19 +++++++++++++++++++ 6 files changed, 45 insertions(+), 4 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py index 125edede25..1e170975a3 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py @@ -44,6 +44,12 @@ def __init__( start_func=start_trans_process_func, up_status_in_queue=up_status_in_queue, ) + + for trans_process in self.kv_trans_processes: + if not trans_process.wait_until_ready(): + raise RuntimeError(f"KV trans module for device {trans_process.device_id} failed to initialize") + + for trans_process in self.kv_trans_processes: threading.Thread(target=self.task_ret_handle_loop, args=(trans_process,), daemon=True).start() # 通过 io buffer 将命令写入到推理进程中 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py index 121a528a43..fcfbaf1936 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py @@ -21,8 +21,11 @@ def start_decode_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - event.wait() - assert proc.is_alive() + while not event.wait(timeout=1): + if not proc.is_alive(): + raise RuntimeError("decode kv move manager process failed during initialization") + if not proc.is_alive(): + raise RuntimeError("decode kv move manager process exited during initialization") logger.info("decode kv move manager process started") return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 2a69fdc523..2ad8820e35 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -47,6 +47,7 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): + module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -100,11 +101,15 @@ def _init_env( up_status_in_queue=up_status_in_queue, ) assert manager is not None + task_out_queue.put("module_ready") + module_ready = True while True: time.sleep(100) except Exception as e: + if not module_ready: + task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py index b23a5c4141..e09ba79c95 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py @@ -20,8 +20,11 @@ def start_prefill_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - event.wait() - assert proc.is_alive() + while not event.wait(timeout=1): + if not proc.is_alive(): + raise RuntimeError("prefill kv move manager process failed during initialization") + if not proc.is_alive(): + raise RuntimeError("prefill kv move manager process exited during initialization") logger.info("prefill kv move manager process started") return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index 9687e6f612..35f2d23cf0 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -41,6 +41,7 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): + module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -74,11 +75,15 @@ def _init_env( mem_managers=mem_managers, ) assert manager is not None + task_out_queue.put("module_ready") + module_ready = True while True: time.sleep(100) except Exception as e: + if not module_ready: + task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py index 073ecf23d2..84173a44ba 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py @@ -1,3 +1,4 @@ +import queue import threading import psutil import torch.multiprocessing as mp @@ -58,5 +59,23 @@ def is_trans_process_health(self): except: return False + def wait_until_ready(self): + for _ in range(600): + try: + status = self.task_out_queue.get(timeout=1) + except queue.Empty: + if not self.process.is_alive(): + logger.error(f"KV trans process for device {self.device_id} exited during initialization") + return False + continue + + if status != "module_ready": + logger.error(f"KV trans module for device {self.device_id} failed to initialize: {status}") + return False + return True + + logger.error(f"KV trans module for device {self.device_id} initialization timed out") + return False + def killself(self): self.process.kill() From 7d5a94a11ad0c4336657920372fae6b1b15472f2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 7 Sep 2026 06:28:30 +0000 Subject: [PATCH 165/214] improve intra-node DP load balancing for PD serving --- .../router/req_queue/dp_balancer/__init__.py | 2 +- .../req_queue/dp_balancer/cache_aware.py | 31 +++++++++++++------ 2 files changed, 22 insertions(+), 11 deletions(-) diff --git a/lightllm/server/router/req_queue/dp_balancer/__init__.py b/lightllm/server/router/req_queue/dp_balancer/__init__.py index 24db8fcd5f..8271236899 100644 --- a/lightllm/server/router/req_queue/dp_balancer/__init__.py +++ b/lightllm/server/router/req_queue/dp_balancer/__init__.py @@ -13,6 +13,6 @@ def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): elif args.dp_balancer == "cache_aware": if args.disable_dynamic_prompt_cache: raise ValueError("cache_aware DP balancing requires dynamic prompt cache") - return DpCacheAwareBalancer(dp_size_in_node, inner_queues) + return DpCacheAwareBalancer(dp_size_in_node, inner_queues, run_mode=args.run_mode) else: raise ValueError(f"Invalid dp balancer: {args.dp_balancer}") diff --git a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py index bbc24c0ff6..6a59270624 100644 --- a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py +++ b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py @@ -2,7 +2,7 @@ The router owns this heuristic index. It records dispatch history rather than querying the infer processes' radix trees, so stale entries can only affect placement, not KV -cache correctness. +cache correctness. Prefill load uses remaining prompt tokens; other modes count requests. """ from __future__ import annotations @@ -99,9 +99,11 @@ def __init__( self, dp_size_in_node: int, inner_queues: List[BaseQueue], + run_mode: str, config: Optional[DpCacheAwareConfig] = None, ) -> None: super().__init__(dp_size_in_node, inner_queues) + self.run_mode = run_mode self.config = config or DpCacheAwareConfig() self.prefix_cache = TokenPrefixCache( block_size=self.config.block_size, @@ -113,16 +115,26 @@ def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: Lis if not reqs_waiting_for_dp_index: return - current_load_per_dp = [0 for _ in range(self.dp_size_in_node)] - if current_batch is not None: - current_load_per_dp = current_batch.get_all_dp_req_num() - total_load_per_dp = [ - current_load_per_dp[dp_index] + len(self.inner_queues[dp_index].waiting_req_list) - for dp_index in range(self.dp_size_in_node) - ] + if self.run_mode == "prefill": + # Queued requests have not matched real KV yet; dispatch history is only a routing hint. + total_load_per_dp = [sum(req.input_len for req in queue.waiting_req_list) for queue in self.inner_queues] + if current_batch is not None: + for req in current_batch.reqs: + total_load_per_dp[req.sample_params.suggested_dp_index] += max( + 0, req.input_len - req.shm_cur_kv_len + ) + else: + current_load_per_dp = [0 for _ in range(self.dp_size_in_node)] + if current_batch is not None: + current_load_per_dp = current_batch.get_all_dp_req_num() + total_load_per_dp = [ + current_load_per_dp[dp_index] + len(self.inner_queues[dp_index].waiting_req_list) + for dp_index in range(self.dp_size_in_node) + ] for req_group in reqs_waiting_for_dp_index: first_req = req_group[0] + group_load = sum(req.input_len for req in req_group) if self.run_mode == "prefill" else len(req_group) linked_prompt_ids = False if not hasattr(first_req, "shm_prompt_ids"): first_req.link_prompt_ids_shm_array() @@ -157,7 +169,6 @@ def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: Lis if cache_dp_index is None: selected_dp_index = least_loaded_dp_index else: - group_load = len(req_group) cache_projected_load = total_load_per_dp[cache_dp_index] + group_load least_projected_load = total_load_per_dp[least_loaded_dp_index] + group_load if cache_projected_load > least_projected_load * self.config.balance_rel_threshold: @@ -168,7 +179,7 @@ def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: Lis for req in req_group: req.sample_params.suggested_dp_index = selected_dp_index self.inner_queues[selected_dp_index].extend(req_group) - total_load_per_dp[selected_dp_index] += len(req_group) + total_load_per_dp[selected_dp_index] += group_load insert_start_index = 0 if cache_dp_index == selected_dp_index: insert_start_index = (matched_token_count + self.config.block_size - 1) // self.config.block_size From 3d910fe7a6001f1494b9cb52f880ff68d6ee4db5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 8 Sep 2026 07:04:50 +0000 Subject: [PATCH 166/214] avoid premature KV transfer aborts and surface generation failures --- lightllm/server/api_http.py | 13 ++++++++++++- lightllm/server/api_openai.py | 10 +++++++++- .../pd/decode_node_impl/decode_trans_process.py | 11 ++++------- lightllm/utils/error_utils.py | 4 ++++ 4 files changed, 29 insertions(+), 9 deletions(-) diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index a2a6036f28..c2388e691a 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -46,7 +46,13 @@ from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.log_utils import init_logger -from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError +from lightllm.utils.error_utils import ( + ClientDisconnected, + GenerationError, + InvalidRequestError, + SERVER_BUSY_MESSAGE, + ServerBusyError, +) from lightllm.server.metrics.manager import MetricClient from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args @@ -181,6 +187,11 @@ async def invalid_request_exception_handler(request: Request, exc: InvalidReques return create_error_response(HTTPStatus.BAD_REQUEST, str(exc)) +@app.exception_handler(GenerationError) +async def generation_exception_handler(request: Request, exc: GenerationError) -> JSONResponse: + return create_error_response(HTTPStatus.INTERNAL_SERVER_ERROR, str(exc)) + + @app.get("/liveness") @app.post("/liveness") def liveness(): diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index 6cd5d31c63..62b300ea85 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -30,7 +30,13 @@ from .httpserver_for_pd_master.manager import HttpServerManagerForPDMaster from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args, get_lightllm_websocket_max_message_size -from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError +from lightllm.utils.error_utils import ( + ClientDisconnected, + GenerationError, + InvalidRequestError, + SERVER_BUSY_MESSAGE, + ServerBusyError, +) from lightllm.utils.log_utils import init_logger from lightllm.server.metrics.manager import MetricClient @@ -599,6 +605,8 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: delta = request_output current_finish_reason = finish_status.get_finish_reason() + if current_finish_reason == "error" and completion_tokens == 1: + raise GenerationError("Generation failed before producing output") # Emit the initial role-only chunk once per choice, as required by the # OpenAI SSE spec: role appears only in the first delta with content="". diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 2ad8820e35..c1ec78835c 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -220,17 +220,12 @@ def dispatch_task_loop(self): trans_task_group: PDChunckedTransTaskGroup = self.recv_task_group_queue.get() with self.waiting_dict_lock: - has_waiting_task = False for task in trans_task_group.task_list: if task.need_transfer_page(): self.waiting_dict[task.get_key()] = task - has_waiting_task = True else: task.start_trans_time = time.time() self.success_queue.put((None, None, task)) - if has_waiting_task: - request_id = trans_task_group.task_list[0].request_id - self.request_last_progress_time[request_id] = time.time() # up status task = trans_task_group.task_list[0] @@ -378,8 +373,10 @@ def _check_tasks_time_out(self): now = time.time() for key, trans_task in list(self.waiting_dict.items()): if trans_task.start_trans_time is None: - request_last_progress = self.request_last_progress_time[trans_task.request_id] - is_timeout = now - request_last_progress > trans_task.time_out_secs + request_last_progress = self.request_last_progress_time.get(trans_task.request_id) + is_timeout = ( + request_last_progress is not None and now - request_last_progress > trans_task.time_out_secs + ) else: is_timeout = trans_task.time_out() if is_timeout: diff --git a/lightllm/utils/error_utils.py b/lightllm/utils/error_utils.py index acf76ad54d..25d7744e50 100644 --- a/lightllm/utils/error_utils.py +++ b/lightllm/utils/error_utils.py @@ -10,6 +10,10 @@ class InvalidRequestError(ValueError): """Request validation failed before generation started.""" +class GenerationError(Exception): + """Generation stopped because of an internal server failure.""" + + class ServerBusyError(Exception): """Custom exception for server busy/overload situations""" From 6b1a28c28ac0e58fa87025b107d2bd09d0ae936f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 8 Sep 2026 14:18:20 +0000 Subject: [PATCH 167/214] fix(vision): map virtual image token IDs before tokenizer decode --- lightllm/models/deepseek_v4/model.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 3e2c47500b..e18e28c17e 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -438,6 +438,15 @@ def encode(self, prompt, multimodal_params=None, **kwargs): return input_ids + def decode(self, token_ids, **kwargs): + if self.has_vision: + # Prompt image embeddings use out-of-vocabulary cache IDs that the HF tokenizer cannot decode. + vocab_size = self.model_config["vocab_size"] + token_ids = [ + self.image_token_id if int(token_id) >= vocab_size else int(token_id) for token_id in token_ids + ] + return self.tokenizer.decode(token_ids, **kwargs) + def _get_encoding_module(self): if self._encoding_module is not None: return self._encoding_module From 7722b82eb210136bd69cd5116ced688e3b8b6424 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Wed, 5 Aug 2026 17:53:26 +0800 Subject: [PATCH 168/214] feat: add EPLB --- docs/CN/source/tutorial/api_server_args.rst | 11 - docs/EN/source/tutorial/api_server_args.rst | 11 - lightllm/common/basemodel/basemodel.py | 9 +- .../layer_infer/cache_tensor_manager.py | 7 +- .../meta_weights/fused_moe/ep_redundancy.py | 195 - .../meta_weights/fused_moe/eplb_placement.py | 506 +++ .../fused_moe/expert_parallel_state.py | 56 + .../fused_moe/fused_moe_weight.py | 107 +- .../meta_weights/fused_moe/impl/__init__.py | 35 +- .../meta_weights/fused_moe/impl/base_impl.py | 91 +- .../fused_moe/impl/deepgemm_impl.py | 215 +- .../fused_moe/impl/marlin_impl.py | 4 + .../fused_moe/impl/triton_impl.py | 99 +- .../triton_kernel/fused_moe/grouped_topk.py | 252 ++ .../triton_kernel/fused_moe/topk_select.py | 9 - .../redundancy_topk_ids_repair.py | 111 - lightllm/common/eplb_utils.py | 16 + .../kv_cache_mem_manager/mem_manager.py | 27 +- lightllm/distributed/communication_op.py | 46 +- lightllm/server/api_cli.py | 15 +- lightllm/server/api_start.py | 10 + lightllm/server/core/objs/start_args_type.py | 4 +- .../model_infer/mode_backend/base_backend.py | 12 +- .../mode_backend/chunked_prefill/impl.py | 5 + .../mode_backend/dp_backend/impl.py | 5 + .../model_infer/mode_backend/eplb_manager.py | 550 +++ .../model_infer/mode_backend/eplb_transfer.py | 775 ++++ .../mode_backend/redundancy_expert_manager.py | 158 - .../server/router/model_infer/model_rpc.py | 8 - lightllm/utils/envs_utils.py | 74 +- lightllm/utils/profile_max_tokens.py | 8 + .../test_redundancy_expert_config.json | 180 - .../test_redundancy_topk_ids_repair.py | 151 - unit_tests/common/fused_moe/test_eplb.py | 3660 +++++++++++++++++ .../fused_moe/test_eplb_transfer_gpu.py | 444 ++ .../models/deepseek_v4/test_memory_profile.py | 65 + unit_tests/server/test_api_start_eplb.py | 64 + 37 files changed, 6825 insertions(+), 1170 deletions(-) delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py delete mode 100644 lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py create mode 100644 lightllm/common/eplb_utils.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_manager.py create mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_transfer.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py delete mode 100644 test/advanced_config/redundancy_expert/test_redundancy_expert_config.json delete mode 100644 unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py create mode 100644 unit_tests/common/fused_moe/test_eplb.py create mode 100644 unit_tests/common/fused_moe/test_eplb_transfer_gpu.py create mode 100644 unit_tests/models/deepseek_v4/test_memory_profile.py create mode 100644 unit_tests/server/test_api_start_eplb.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 746827b102..0e10932916 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -666,17 +666,6 @@ MTP 多预测参数 增加此值允许更多预测,但确保模型与指定的步数兼容。 目前 deepseekv3/r1 模型仅支持 1 步 -DeepSeek 冗余专家参数 ---------------------- - -.. option:: --ep_redundancy_expert_config_path - - 冗余专家配置的路径。可用于 deepseekv3 模型。 - -.. option:: --auto_update_redundancy_expert - - 是否通过在线专家使用计数器为 deepseekv3 模型更新冗余专家。 - 监控和日志参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 81b3bb7523..7efa249eaf 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -673,17 +673,6 @@ MTP Multi-Prediction Parameters Increasing this value allows more predictions, but ensure the model is compatible with the specified number of steps. Currently deepseekv3/r1 models only support 1 step -DeepSeek Redundant Expert Parameters ------------------------------------- - -.. option:: --ep_redundancy_expert_config_path - - Path to redundant expert configuration. Can be used for deepseekv3 models. - -.. option:: --auto_update_redundancy_expert - - Whether to update redundant experts for deepseekv3 models through online expert usage counters. - Monitoring and Logging Parameters --------------------------------- diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 3025cbf269..88e50a0b5e 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -73,6 +73,7 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() + self.eplb_manager = None self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] @@ -381,13 +382,15 @@ def forward(self, model_input: ModelInput): if model_input.is_prefill: model_output = self._prefill(model_input) - self._record_prefill_ep_balance() + self._after_prefill() return model_output return self._decode(model_input) - def _record_prefill_ep_balance(self): + def _after_prefill(self): if self.ep_balance_monitor is not None: self.ep_balance_monitor.record_prefill_round() + if self.eplb_manager is not None: + self.eplb_manager.step() def _create_inferstate(self, model_input: ModelInput, microbatch_index: int = 0): infer_state = self.infer_state_class() @@ -876,7 +879,7 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event - self._record_prefill_ep_balance() + self._after_prefill() return model_output0, model_output1 @torch.no_grad() diff --git a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py index 8bcf99b992..13906d0d8a 100644 --- a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py +++ b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py @@ -22,7 +22,12 @@ def custom_del(self: torch.Tensor): if hasattr(self, "storage_weak_ptr"): storage_weak_ptr = self.storage_weak_ptr else: - storage_weak_ptr = self.untyped_storage()._weak_ref() + try: + storage_weak_ptr = self.untyped_storage()._weak_ref() + except RuntimeError: + # Some tensor implementations, including UndefinedTensorImpl, + # have no backing storage. Their destructor must stay silent. + return UntypedStorage._free_weak_ref(storage_weak_ptr) if storage_weak_ptr in g_cache_manager.ptr_to_bufnode: g_cache_manager.changed_ptr.add(storage_weak_ptr) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py deleted file mode 100644 index 749400c8d8..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py +++ /dev/null @@ -1,195 +0,0 @@ -import numpy as np -import torch -from .fused_moe_weight import FusedMoeWeight -from lightllm.utils.log_utils import init_logger -from typing import Dict - -logger = init_logger(__name__) - - -class FusedMoeWeightEPAutoRedundancy: - def __init__( - self, - ep_fused_moe_weight: FusedMoeWeight, - ) -> None: - super().__init__() - self._ep_w = ep_fused_moe_weight - self.redundancy_expert_num = self._ep_w.redundancy_expert_num - - def clear_counter(self): - self._ep_w.routed_expert_counter_tensor.fill_(0) - return - - def prepare_redundancy_experts( - self, - ): - expert_counter = self._ep_w.routed_expert_counter_tensor.detach().cpu().numpy() - logger.info( - f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" - f" expert_counter: {expert_counter}" - ) - self._ep_w.routed_expert_counter_tensor.fill_(0) - ep_n_routed_experts = self._ep_w.n_routed_experts // self._ep_w.global_world_size - start_expert_id = ep_n_routed_experts * self._ep_w.global_rank_ - no_redundancy_expert_ids = list(range(start_expert_id, start_expert_id + ep_n_routed_experts)) - - # 统计 0 rank 上的全局 topk 冗余信息,帮助导出一份全局可用的静态使用的冗余专家静态配置。 - if self._ep_w.global_rank_ == 0: - # int(e) for serialization, int64 can not be serialized by json.dump. - topk_redundancy_expert_ids = list(int(e) for e in np.argsort(expert_counter)[-self.redundancy_expert_num :]) - else: - topk_redundancy_expert_ids = None - - # 不要选中当前已经存在的非冗余专家作为冗余专家 - expert_counter[no_redundancy_expert_ids] = 0 - - self.redundancy_expert_ids = list(np.argsort(expert_counter)[-self.redundancy_expert_num :]) - logger.info( - f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" - f" new select redundancy_expert_ids : {self.redundancy_expert_ids}" - ) - - # 准备加载过度变量。 - self.experts_up_projs = [None] * self.redundancy_expert_num - self.experts_gate_projs = [None] * self.redundancy_expert_num - self.experts_up_proj_scales = [None] * self.redundancy_expert_num - self.experts_gate_proj_scales = [None] * self.redundancy_expert_num - self.w2_list = [None] * self.redundancy_expert_num - self.w2_scale_list = [None] * self.redundancy_expert_num - self.w13 = [None, None] # weight, weight_scale - self.w2 = [None, None] # weight, weight_scale - return topk_redundancy_expert_ids - - def load_hf_weights(self, weights): - # 加载冗余专家的权重参数 - for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): - i_experts = redundant_expert_id - w1_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.weight" - w2_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.weight" - w3_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.weight" - if w1_weight in weights: - self.experts_gate_projs[i] = weights[w1_weight] - if w3_weight in weights: - self.experts_up_projs[i] = weights[w3_weight] - if w2_weight in weights: - self.w2_list[i] = weights[w2_weight] - - self._load_weight_scale(weights) - self._fuse() - - def _fuse(self): - self._fuse_weight_scale() - - with self._ep_w.lock: - if ( - hasattr(self, "experts_up_projs") - and None not in self.experts_up_projs - and None not in self.experts_gate_projs - and None not in self.w2_list - ): - gate_out_dim, gate_in_dim = self.experts_gate_projs[0].shape - up_out_dim, up_in_dim = self.experts_up_projs[0].shape - assert gate_in_dim == up_in_dim - dtype = self.experts_gate_projs[0].dtype - total_expert_num = self.redundancy_expert_num - - w13 = torch.empty((total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu") - - for i_experts in range(self.redundancy_expert_num): - w13[i_experts, 0:gate_out_dim:, :] = self.experts_gate_projs[i_experts] - w13[i_experts, gate_out_dim:, :] = self.experts_up_projs[i_experts] - - inter_shape, hidden_size = self.w2_list[0].shape[0], self.w2_list[0].shape[1] - w2 = torch._utils._flatten_dense_tensors(self.w2_list).view(len(self.w2_list), inter_shape, hidden_size) - if self._ep_w.quant_method._check_weight_need_quanted(weight=w13): - w13_pack, _ = self._ep_w.quant_method.create_moe_weight( - out_dims=[gate_out_dim + up_out_dim], - in_dim=1, - dtype=self._ep_w.data_type_, - device_id=self._ep_w.device_id_, - num_experts=self.redundancy_expert_num, - ) - self._ep_w.quant_method.quantize(w13, w13_pack) - w2_pack, _ = self._ep_w.quant_method.create_moe_weight( - out_dims=[inter_shape], - in_dim=hidden_size, - dtype=self._ep_w.data_type_, - device_id=self._ep_w.device_id_, - num_experts=self.redundancy_expert_num, - ) - self._ep_w.quant_method.quantize(w2, w2_pack) - - self.w13[0] = w13_pack.weight - self.w13[1] = w13_pack.weight_scale - self.w2[0] = w2_pack.weight - self.w2[1] = w2_pack.weight_scale - else: - self.w13[0] = w13 - self.w2[0] = w2 - delattr(self, "w2_list") - delattr(self, "experts_up_projs") - delattr(self, "experts_gate_projs") - - def _fuse_weight_scale(self): - with self._ep_w.lock: - if ( - hasattr(self, "experts_up_proj_scales") - and None not in self.experts_up_proj_scales - and None not in self.experts_gate_proj_scales - and None not in self.w2_scale_list - ): - gate_out_dim, gate_in_dim = self.experts_gate_proj_scales[0].shape - up_out_dim, up_in_dim = self.experts_up_proj_scales[0].shape - assert gate_in_dim == up_in_dim - dtype = self.experts_gate_proj_scales[0].dtype - total_expert_num = self.redundancy_expert_num - w13_scale = torch.empty( - (total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu" - ) - for i_experts in range(self.redundancy_expert_num): - w13_scale[i_experts, 0:gate_out_dim:, :] = self.experts_gate_proj_scales[i_experts] - w13_scale[i_experts, gate_out_dim:, :] = self.experts_up_proj_scales[i_experts] - - inter_shape, hidden_size = self.w2_scale_list[0].shape[0], self.w2_scale_list[0].shape[1] - w2_scale = torch._utils._flatten_dense_tensors(self.w2_scale_list).view( - len(self.w2_scale_list), inter_shape, hidden_size - ) - self.w13[1] = w13_scale - self.w2[1] = w2_scale - delattr(self, "w2_scale_list") - delattr(self, "experts_up_proj_scales") - delattr(self, "experts_gate_proj_scales") - - def _load_weight_scale(self, weights: Dict[str, torch.Tensor]) -> None: - # 加载冗余专家的scale参数 - for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): - i_experts = redundant_expert_id - weight_scale_suffix = self._ep_w.quant_method.weight_scale_suffix - w1_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.{weight_scale_suffix}" - w2_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.{weight_scale_suffix}" - w3_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.{weight_scale_suffix}" - if w1_scale in weights: - self.experts_gate_proj_scales[i] = weights[w1_scale] - if w3_scale in weights: - self.experts_up_proj_scales[i] = weights[w3_scale] - if w2_scale in weights: - self.w2_scale_list[i] = weights[w2_scale] - - def commit(self): - for index, dest_tensor in enumerate([self._ep_w.w13.weight, self._ep_w.w13.weight_scale]): - if dest_tensor is not None: - assert isinstance( - dest_tensor, torch.Tensor - ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" - dest_tensor[-self.redundancy_expert_num :, :, :] = self.w13[index][:, :, :] - - for index, dest_tensor in enumerate([self._ep_w.w2.weight, self._ep_w.w2.weight_scale]): - if dest_tensor is not None: - assert isinstance( - dest_tensor, torch.Tensor - ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" - dest_tensor[-self.redundancy_expert_num :, :, :] = self.w2[index][:, :, :] - - self._ep_w.redundancy_expert_ids_tensor.copy_( - torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cpu") - ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py new file mode 100644 index 0000000000..d550355b9f --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -0,0 +1,506 @@ +from dataclasses import dataclass +from functools import lru_cache +from typing import Dict, Tuple +import torch + + +def build_initial_redundant_expert_ids( + num_logical_experts: int, + num_ranks: int, + num_redundant_experts_per_rank: int, +) -> torch.Tensor: + """Build a deterministic initial placement without local duplicates.""" + assert num_logical_experts % num_ranks == 0 + num_experts_per_rank = num_logical_experts // num_ranks + assert 0 < num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank + + # 初始化结果确定,不依赖随机数。 + # 每个 rank 不会复制自己原本拥有的 expert。 + # 同一个 rank 的冗余槽位不会重复。 + # 最后一个 rank 通过取模自然回绕。 + rank_offsets = torch.arange(1, num_ranks + 1, dtype=torch.int64)[:, None] * num_experts_per_rank + expert_offsets = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64) + return (rank_offsets + expert_offsets) % num_logical_experts + + +def build_logical_to_physical_map( + redundant_expert_ids: torch.Tensor, # 冗余布局,shape 为 [num_ranks, num_redundant_experts_per_rank]。 + num_logical_experts: int, # 逻辑 expert 的总数。 + source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 + node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 +) -> Tuple[ + torch.Tensor, torch.Tensor +]: # logical_to_physical [num_logical_experts, num_ranks], replica_counts [num_logical_experts] + """构建单层逻辑 expert 到物理副本的映射""" + + logical_to_physical, replica_counts = build_logical_to_physical_maps_for_layers( + redundant_expert_ids.unsqueeze(0), + num_logical_experts, + source_rank=source_rank, + node_world_size=node_world_size, + ) + return logical_to_physical.squeeze(0), replica_counts.squeeze(0) + + +def build_logical_to_physical_maps_for_layers( + redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] + num_logical_experts: int, # 逻辑 expert 的总数。 + source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 + node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 +) -> Tuple[ + torch.Tensor, # logical_to_physical, shape [num_layers, num_logical_experts, num_ranks] + torch.Tensor, # replica_counts, shape [num_layers, num_logical_experts] +]: + """Build stable CPU int32 maps for the supplied layers without modifying the input.""" + if redundant_expert_ids_by_layer.ndim != 3: + raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") + num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] + assert num_logical_experts % num_ranks == 0 + layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) + logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) + if source_rank is None: + return logical_to_physical, replica_counts + + assert node_world_size is not None + replica_positions = torch.arange(num_ranks, dtype=torch.int64) + compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( + logical_to_physical, + replica_counts, + source_rank=source_rank, + node_world_size=node_world_size, + num_physical_experts_per_rank=layout.num_physical_experts_per_rank, + replica_positions=replica_positions, + ) + return ( + _rotate_selected_replicas( + compact_maps_by_layer, + selected_counts_by_layer, + source_rank=source_rank, + replica_positions=replica_positions, + ), + selected_counts_by_layer, + ) + + +def select_improving_placements( + expert_load: torch.Tensor, + current_placement: torch.Tensor, + candidate_placement: torch.Tensor, + *, + rebalance_gain_threshold: float, + expert_alignment: int | None = None, + node_world_size: int | None = None, +) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, float | int], torch.Tensor, torch.Tensor]: + """Select better layers and return current/final rank loads without re-estimation.""" + if not 0.0 <= rebalance_gain_threshold <= 1.0: + raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") + assert current_placement.shape == candidate_placement.shape + current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) + candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) + current_critical = current_rank_load.max(dim=-1).values + candidate_critical = candidate_rank_load.max(dim=-1).values + if current_rank_load.ndim == 3: + current_critical = current_critical.sum(dim=0) + candidate_critical = candidate_critical.sum(dim=0) + # Each changed layer must reduce its own critical load. All selected + # changes must then collectively meet the configured model-level + # critical-load reduction threshold, avoiding low-gain migrations. + improved = candidate_critical < current_critical + selected = current_placement.clone() + selected[improved] = candidate_placement[improved] + selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) + model_current_critical = current_critical.sum() + model_current_mean = current_rank_load.mean(dim=-1).sum() + model_selected_critical = selected_rank_load.max(dim=-1).values.sum() + model_selected_mean = selected_rank_load.mean(dim=-1).sum() + model_ratio = model_current_critical / model_current_mean.clamp_min(1.0) + candidate_model_ratio = model_selected_critical / model_selected_mean.clamp_min(1.0) + candidate_rebalance_gain = (model_current_critical - model_selected_critical) / model_current_critical.clamp_min( + 1.0 + ) + metrics = { + "model_imbalance_ratio": float(model_ratio.item()), + "candidate_model_imbalance_ratio": float(candidate_model_ratio.item()), + "candidate_rebalance_gain": float(candidate_rebalance_gain.item()), + "candidate_changed_layer_count": int(improved.sum().item()), + } + if candidate_rebalance_gain >= rebalance_gain_threshold: + return selected, improved, metrics, current_rank_load, selected_rank_load + return ( + current_placement.clone(), + torch.zeros_like(improved), + metrics, + current_rank_load, + current_rank_load, + ) + + +def plan_redundant_experts( + expert_load: torch.Tensor, + num_ranks: int, + num_redundant_experts_per_rank: int, + expert_alignment: int | None = None, + node_world_size: int | None = None, + current_placement: torch.Tensor | None = None, + stickiness: float = 0.0, +) -> torch.Tensor: + """Plan replicas using source-node-local copies, with global fallback. + + With ``current_placement`` and a positive ``stickiness``, a candidate that + keeps an expert on its current rank receives a bonus of + ``stickiness * mean per-layer expert load``. This preserves rank + membership, not a particular redundant physical slot; target slots are + canonicalized against the current live rows before transfer and metadata + publication. A rank membership only changes when the move improves the + critical-load objective by more than that margin. + Without them the planning is bit-identical to the legacy behavior. + """ + assert expert_load.ndim in (2, 3, 4) + if expert_alignment is not None: + assert expert_alignment > 0 + use_legacy_topology_preference = expert_load.ndim < 4 + legacy_node_world_size = node_world_size if use_legacy_topology_preference else None + source_load, _squeeze_sample, node_world_size = _as_source_node_load(expert_load, num_ranks, node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + assert num_logical_experts % num_ranks == 0 + assert num_redundant_experts_per_rank > 0 + num_experts_per_rank = num_logical_experts // num_ranks + num_redundant = num_ranks * num_redundant_experts_per_rank + assert num_redundant <= num_logical_experts * (num_ranks - 1) + + load = source_load.to(dtype=torch.float64, device="cpu") + placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) + owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank + if current_placement is not None: + assert tuple(current_placement.shape) == ( + num_layers, + num_ranks, + num_redundant_experts_per_rank, + ) + current_locations = _expert_locations(current_placement, num_logical_experts) + stickiness_scale = load.sum(dim=(0, 2, 3)) / num_logical_experts + else: + current_locations = None + stickiness_scale = None + + locations = _expert_locations(placement, num_logical_experts) + expert_rank = _expert_rank_load_all(load, locations, num_nodes, node_world_size, expert_alignment) + rank_load = expert_rank.sum(dim=2) + remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) + layer_indices = torch.arange(num_layers, dtype=torch.int64) + expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) + rank_nodes = ( + torch.arange(num_ranks, dtype=torch.int64) // legacy_node_world_size + if legacy_node_world_size is not None and legacy_node_world_size < num_ranks + else None + ) + + # Every iteration fills one slot per layer. Candidate expert evaluation + # is vectorized across all layers and logical experts, which keeps large + # GLM/Qwen planning comfortably on the CPU fast path. + for _ in range(num_redundant): + rank_order = torch.argsort(rank_load.sum(dim=0), dim=1, stable=True) + target_ranks = torch.full((num_layers,), -1, dtype=torch.int64) + legal = torch.zeros((num_layers, num_logical_experts), dtype=torch.bool) + for layer in range(num_layers): + for target_rank in rank_order[layer].tolist(): + if remaining_slots[layer, target_rank] == 0: + continue + candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] + # Legacy 2D/3D callers have no source-node axis. Retain the + # previous topology preference for that compatibility path; + # node-aware [S,L,N,E] planning uses only the exact load + # objective below. + if rank_nodes is not None: + existing_on_target_node = locations[layer, :, rank_nodes == rank_nodes[target_rank]].any(dim=1) + new_node_legal = candidate_legal & ~existing_on_target_node + if torch.any(new_node_legal): + candidate_legal = new_node_legal + if torch.any(candidate_legal): + target_ranks[layer] = target_rank + legal[layer] = candidate_legal + break + if torch.any(target_ranks < 0): + raise RuntimeError("EPLB planner found no valid redundant expert placement") + + candidate_locations = locations.clone() + candidate_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] = True + candidate_expert_rank = _expert_rank_load_all( + load, candidate_locations, num_nodes, node_world_size, expert_alignment + ) + candidate_rank_load = rank_load[:, :, None, :] - expert_rank + candidate_expert_rank + critical = candidate_rank_load.max(dim=3).values.sum(dim=0) + critical.masked_fill_(~legal, torch.inf) + if current_locations is not None: + # An expert already held by the target rank is retained unless + # another candidate beats it by more than the stickiness margin. + # This is rank membership, not physical-slot stickiness. Masked + # (inf) candidates stay masked: inf - x == inf. + keep = current_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] + critical = critical - stickiness * stickiness_scale[:, None] * keep + selected_experts = critical.argmin(dim=1) + if torch.isinf(critical[layer_indices, selected_experts]).any(): + raise RuntimeError("EPLB planner found no valid redundant expert placement") + + slots = num_redundant_experts_per_rank - remaining_slots[layer_indices, target_ranks] + placement[layer_indices, target_ranks, slots] = selected_experts + selected_next = candidate_expert_rank[:, layer_indices, selected_experts] + selected_old = expert_rank[:, layer_indices, selected_experts] + rank_load += selected_next - selected_old + expert_rank[:, layer_indices, selected_experts] = selected_next + locations[layer_indices, selected_experts, target_ranks] = True + remaining_slots[layer_indices, target_ranks] -= 1 + + assert torch.all(placement >= 0) + return placement + + +@dataclass(frozen=True, eq=False) +class _PhysicalExpertLayout: + """进程内按拓扑复用的只读物理 expert 布局;其中 Tensor 不得原地修改。""" + + num_logical_experts: int + num_ranks: int + num_physical_experts_per_rank: int + primary_physical_ids: torch.Tensor + redundant_physical_ids: torch.Tensor + + +@lru_cache(maxsize=8) +def _get_physical_expert_layout( + num_logical_experts: int, + num_ranks: int, + num_redundant_experts_per_rank: int, +) -> _PhysicalExpertLayout: + """返回按静态拓扑缓存的只读 CPU 物理 expert ID。""" + num_experts_per_rank = num_logical_experts // num_ranks + num_physical_experts_per_rank = num_experts_per_rank + num_redundant_experts_per_rank + expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) + primary_physical_ids = ( + (expert_ids // num_experts_per_rank) * num_physical_experts_per_rank + expert_ids % num_experts_per_rank + ).to(torch.int32) + ranks = torch.arange(num_ranks, dtype=torch.int64).repeat_interleave(num_redundant_experts_per_rank) + slots = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64).repeat(num_ranks) + redundant_physical_ids = (ranks * num_physical_experts_per_rank + num_experts_per_rank + slots).to(torch.int32) + return _PhysicalExpertLayout( + num_logical_experts=num_logical_experts, + num_ranks=num_ranks, + num_physical_experts_per_rank=num_physical_experts_per_rank, + primary_physical_ids=primary_physical_ids, + redundant_physical_ids=redundant_physical_ids, + ) + + +def _build_global_replica_maps_for_layers( + redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] + layout: _PhysicalExpertLayout, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Build stable global maps with primary copies first and unused slots set to ``-1``.""" + + num_layers = redundant_expert_ids_by_layer.shape[0] + num_logical_experts = layout.num_logical_experts + max_replicas = layout.num_ranks + redundant_ids = redundant_expert_ids_by_layer.to(dtype=torch.int64, device="cpu") + logical_to_physical = torch.full((num_layers, num_logical_experts, max_replicas), -1, dtype=torch.int32) + logical_to_physical[:, :, 0] = layout.primary_physical_ids + replica_counts = torch.ones((num_layers, num_logical_experts), dtype=torch.int32) + + flat_redundant_ids = redundant_ids.reshape(num_layers, -1) + if not flat_redundant_ids.numel(): + return logical_to_physical, replica_counts + + # 稳定排序保留 rank-major、slot-major 的历史顺序;第 0 列固定为主副本。 + sort_order = torch.argsort(flat_redundant_ids, dim=1, stable=True) + sorted_redundant_ids = flat_redundant_ids.gather(1, sort_order) + flat_positions = torch.arange(flat_redundant_ids.shape[1], dtype=torch.int64).unsqueeze(0) + group_starts = torch.where( + torch.cat( + ( + torch.ones((num_layers, 1), dtype=torch.bool), + sorted_redundant_ids[:, 1:] != sorted_redundant_ids[:, :-1], + ), + dim=1, + ), + flat_positions, + 0, + ) + replica_indices = flat_positions - torch.cummax(group_starts, dim=1).values + 1 + redundant_counts = torch.zeros((num_layers, num_logical_experts), dtype=torch.int32) + redundant_counts.scatter_add_( + 1, + flat_redundant_ids, + torch.ones_like(flat_redundant_ids, dtype=torch.int32), + ) + assert int(redundant_counts.max().item()) < max_replicas, "an expert can have at most one replica per rank" + replica_counts += redundant_counts + + layer_indices = torch.arange(num_layers, dtype=torch.int64).view(-1, 1).expand_as(sort_order) + redundant_physical_ids = layout.redundant_physical_ids.unsqueeze(0).expand_as(sort_order).gather(1, sort_order) + logical_to_physical[layer_indices, sorted_redundant_ids, replica_indices] = redundant_physical_ids + return logical_to_physical, replica_counts + + +def _select_source_node_replicas( + logical_to_physical: torch.Tensor, + replica_counts: torch.Tensor, + *, + source_rank: int, + node_world_size: int, + num_physical_experts_per_rank: int, + replica_positions: torch.Tensor, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Keep source-node replicas when available, otherwise fall back to all stable candidates.""" + num_layers, num_logical_experts, _max_replicas = logical_to_physical.shape + source_node = source_rank // node_world_size + output_positions = replica_positions.view(1, 1, -1) + valid = output_positions < replica_counts.unsqueeze(-1) + local = valid & ( + torch.div( + logical_to_physical, + num_physical_experts_per_rank * node_world_size, + rounding_mode="floor", + ) + == source_node + ) + selected = torch.where(local.any(dim=2, keepdim=True), local, valid) + selected_counts_by_layer = selected.sum(dim=2, dtype=torch.int32) + + compact_maps_by_layer = torch.full_like(logical_to_physical, -1) + selected_positions = selected.cumsum(dim=2) - 1 + layers = torch.arange(num_layers, dtype=torch.int64).view(-1, 1, 1).expand_as(selected) + experts = torch.arange(num_logical_experts, dtype=torch.int64).view(1, -1, 1).expand_as(selected) + compact_maps_by_layer[layers[selected], experts[selected], selected_positions[selected]] = logical_to_physical[ + selected + ] + return compact_maps_by_layer, selected_counts_by_layer + + +def _rotate_selected_replicas( + compact_maps_by_layer: torch.Tensor, + selected_count_by_layer: torch.Tensor, + *, + source_rank: int, + replica_positions: torch.Tensor, +) -> torch.Tensor: + """根据 source_rank 循环调整每个逻辑 expert 的候选副本顺序,让不同源 rank 优先使用不同副本,同时保持候选副本集合和副本数量不变""" + output_positions = replica_positions.view(1, 1, -1) + selected_count64_by_layer = selected_count_by_layer.to(torch.int64).unsqueeze(-1) + source_positions_by_layer = (output_positions + source_rank) % selected_count64_by_layer + maps_by_layer = compact_maps_by_layer.gather(2, source_positions_by_layer) + maps_by_layer.masked_fill_(output_positions >= selected_count64_by_layer, -1) + return maps_by_layer + + +def _estimate_rank_load( + expert_load: torch.Tensor, + redundant_expert_ids: torch.Tensor, + expert_alignment: int | None = None, + node_world_size: int | None = None, +) -> torch.Tensor: + """Estimate runtime source-node-local routing load per physical expert. + + ``expert_load`` accepts the historic ``[layers, experts]`` and + ``[samples, layers, experts]`` forms, which are both one source node, and + the distributed ``[samples, layers, source_nodes, experts]`` form. Source + loads are kept separate until they are assigned to physical replicas, then + combined before applying the per-expert alignment used by DeepEP. + """ + source_load, squeeze_sample, node_world_size = _as_source_node_load( + expert_load, redundant_expert_ids.shape[1], node_world_size + ) + num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers + num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] + assert num_logical_experts % num_ranks == 0 + if expert_alignment is not None: + assert expert_alignment > 0 + + rank_load = _expert_rank_load_all( + source_load, + _expert_locations(redundant_expert_ids, num_logical_experts), + num_nodes, + node_world_size, + expert_alignment, + ).sum(dim=2) + return rank_load.squeeze(0) if squeeze_sample else rank_load + + +def _as_source_node_load( + expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None +) -> Tuple[torch.Tensor, bool, int]: + """Normalize load to ``[samples, layers, source_nodes, experts]``.""" + assert expert_load.ndim in (2, 3, 4) + squeeze_sample = expert_load.ndim == 2 + if expert_load.ndim == 2: + source_load = expert_load.unsqueeze(0).unsqueeze(2) + elif expert_load.ndim == 3: + source_load = expert_load.unsqueeze(2) + else: + source_load = expert_load + num_nodes = source_load.shape[2] + # Historic 2D/3D loads represent one source node containing every rank. + if expert_load.ndim < 4: + return source_load, squeeze_sample, num_ranks + if node_world_size is None: + assert num_ranks % num_nodes == 0 + node_world_size = num_ranks // num_nodes + assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 + assert num_nodes == num_ranks // node_world_size + return source_load, squeeze_sample, node_world_size + + +def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: + """Return ``[layer, logical expert, rank]`` physical-copy occupancy.""" + num_layers, num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape + assert num_logical_experts % num_ranks == 0 + num_experts_per_rank = num_logical_experts // num_ranks + locations = torch.zeros( + (num_layers, num_logical_experts, num_ranks), + dtype=torch.bool, + device=redundant_expert_ids.device, + ) + expert_ids = torch.arange(num_logical_experts, device=locations.device) + owners = expert_ids // num_experts_per_rank + locations[:, expert_ids, owners] = True + layers = torch.arange(num_layers, device=locations.device)[:, None] + ranks = torch.arange(num_ranks, device=locations.device).repeat_interleave(num_redundant_experts_per_rank)[None, :] + redundant_ids = redundant_expert_ids.reshape(num_layers, -1) + valid = redundant_ids >= 0 + if torch.any(valid): + expanded_layers = layers.expand_as(redundant_ids) + expanded_ranks = ranks.expand_as(redundant_ids) + locations[ + expanded_layers[valid], + redundant_ids[valid], + expanded_ranks[valid], + ] = True + return locations + + +def _source_route(slots: torch.Tensor, num_nodes: int, node_world_size: int) -> torch.Tensor: + """Route each source node to its local copies, or all copies as fallback.""" + num_ranks = slots.shape[-1] + assert num_ranks % node_world_size == 0 and num_nodes == num_ranks // node_world_size + rank_nodes = torch.arange(num_ranks, device=slots.device) // node_world_size + source_nodes = torch.arange(num_nodes, device=slots.device) + copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_nodes, *slots.shape[-2:]) + rank_node_shape = (1,) * slots.ndim + (num_ranks,) + source_node_shape = (1,) * (slots.ndim - 2) + (num_nodes, 1, 1) + local = copies & (rank_nodes.reshape(rank_node_shape) == source_nodes.reshape(source_node_shape)) + selected = torch.where(local.any(dim=-1, keepdim=True), local, copies) + return selected.to(torch.float64) / selected.sum(dim=-1, keepdim=True) + + +def _expert_rank_load_all( + source_load: torch.Tensor, + locations: torch.Tensor, + num_nodes: int, + node_world_size: int, + expert_alignment: int | None, +) -> torch.Tensor: + """Return aligned ``[samples, layers, expert, rank]`` contributions.""" + route = _source_route(locations, num_nodes, node_world_size) + physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) + if expert_alignment is not None: + physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment + return physical_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py new file mode 100644 index 0000000000..f66d31a62b --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py @@ -0,0 +1,56 @@ +from contextlib import contextmanager +from contextvars import ContextVar +from dataclasses import dataclass +from typing import Iterator, Optional + +import torch + + +_eplb_model_init_disabled: ContextVar[bool] = ContextVar("eplb_model_init_disabled", default=False) + + +def is_eplb_model_init_disabled() -> bool: + return _eplb_model_init_disabled.get() + + +@contextmanager +def disable_eplb_model_init() -> Iterator[None]: + token = _eplb_model_init_disabled.set(True) + try: + yield + finally: + _eplb_model_init_disabled.reset(token) + + +@dataclass +class EPLBState: + num_redundant_experts_per_rank: int + initial_redundant_expert_ids_by_rank: torch.Tensor + logical_to_physical_map: torch.Tensor + logical_replica_count: torch.Tensor + route_counter: torch.Tensor + recording: bool = False + recorded_sample_count: int = 0 + + def next_sample_index(self) -> int: + if not self.recording: + return 0 + sample_index = self.recorded_sample_count % self.route_counter.shape[0] + self.recorded_sample_count += 1 + return sample_index + + +@dataclass(frozen=True) +class ExpertParallelState: + num_logical_experts: int + world_size: int + eplb: Optional[EPLBState] = None + + @property + def num_primary_experts_per_rank(self) -> int: + return self.num_logical_experts // self.world_size + + @property + def num_total_physical_experts(self) -> int: + num_redundant_experts_per_rank = 0 if self.eplb is None else self.eplb.num_redundant_experts_per_rank + return self.num_logical_experts + self.world_size * num_redundant_experts_per_rank diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 758e09cdfb..3f7c5fe6b6 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -8,11 +8,20 @@ get_col_slice_mixin, SliceMixinTpl, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import select_fuse_moe_impl +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import create_fuse_moe_impl +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, + is_eplb_model_init_disabled, +) from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, + build_logical_to_physical_map, +) from lightllm.common.quantization.quantize_method import QuantizationMethod -from lightllm.utils.envs_utils import get_redundancy_expert_ids, get_redundancy_expert_num, get_env_start_args -from lightllm.utils.dist_utils import get_global_world_size, get_global_rank +from lightllm.utils.envs_utils import get_env_start_args, get_prefill_eplb_step_interval +from lightllm.utils.dist_utils import get_global_world_size, get_global_rank, get_node_world_size from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -56,17 +65,14 @@ def __init__( self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self._init_config(network_config) - self._init_redundancy_expert_params() - self._init_parallel_params() - self.fuse_moe_impl = select_fuse_moe_impl(self.quant_method, self.enable_ep_moe)( + self._init_expert_parallel_state() + self._init_weight_partition() + self.fuse_moe_impl = create_fuse_moe_impl( n_routed_experts=self.n_routed_experts, num_fused_shared_experts=self.num_fused_shared_experts, routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, - redundancy_expert_num=self.redundancy_expert_num, - redundancy_expert_ids_tensor=self.redundancy_expert_ids_tensor, - routed_expert_counter_tensor=self.routed_expert_counter_tensor, - auto_update_redundancy_expert=self.auto_update_redundancy_expert, + expert_parallel_state=self.expert_parallel_state, ) self.lock = threading.Lock() self._moe_weight_finalized = False @@ -81,16 +87,50 @@ def _init_config(self, network_config: Dict[str, Any]): self.routed_scaling_factor = network_config.get("routed_scaling_factor", 1.0) self.scoring_func = network_config.get("scoring_func", "softmax") - def _init_redundancy_expert_params(self): - self.redundancy_expert_num = get_redundancy_expert_num() - self.redundancy_expert_ids = get_redundancy_expert_ids(self.layer_num_) - self.auto_update_redundancy_expert: bool = get_env_start_args().auto_update_redundancy_expert - self.redundancy_expert_ids_tensor = torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cuda") - self.routed_expert_counter_tensor = torch.zeros((self.n_routed_experts,), dtype=torch.int64, device="cuda") - # TODO: find out the reason of failure of deepep when redundancy_expert_num is 1. - assert self.redundancy_expert_num != 1, "redundancy_expert_num can not be 1 for some unknown hang of deepep." + def _init_expert_parallel_state(self): + args = get_env_start_args() + self.expert_parallel_state: Optional[ExpertParallelState] = None + # Initial placement metadata is used only while loading checkpoint rows. + self._initial_redundant_expert_ids = [] + self._initial_redundant_expert_idx_to_local_idx = {} + eplb = None + if args.enable_prefill_eplb and not is_eplb_model_init_disabled(): + num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank + all_initial_ids = build_initial_redundant_expert_ids( + self.n_routed_experts, + self.global_world_size, + num_redundant_experts_per_rank, + ) + self._initial_redundant_expert_ids = all_initial_ids[self.global_rank_].tolist() + logical_to_physical, logical_replica_count = build_logical_to_physical_map( + all_initial_ids, + self.n_routed_experts, + source_rank=self.global_rank_, + node_world_size=get_node_world_size(), + ) + # route_counter 每次 prefill dispatch 记录一行。初始阶段连续采样 + # step_interval 个 manager step,兼顾micro batch overlap的两次 dispatch,因此容量设为 + # 2 * step_interval。稳定阶段复用该环形缓冲区,但只把当前短采样窗口内 + # 实际记录的最近行传给 planner,不复制整个缓冲区。 + eplb = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=all_initial_ids, + logical_to_physical_map=logical_to_physical.cuda(), + logical_replica_count=logical_replica_count.cuda(), + route_counter=torch.zeros( + (2 * get_prefill_eplb_step_interval(), self.n_routed_experts), + dtype=torch.int64, + device="cuda", + ), + ) + if self.enable_ep_moe: + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=self.n_routed_experts, + world_size=self.global_world_size, + eplb=eplb, + ) - def _init_parallel_params(self): + def _init_weight_partition(self): if self.enable_ep_moe: self.tp_rank_ = 0 self.tp_world_size_ = 1 @@ -104,27 +144,26 @@ def _init_parallel_params(self): self.split_inter_size = self.moe_intermediate_size // self.tp_world_size_ if self.enable_ep_moe: assert self.num_fused_shared_experts == 0, "num_fused_shared_experts must be 0 when enable_ep_moe" + eplb = self.expert_parallel_state.eplb + num_redundant_experts_per_rank = 0 if eplb is None else eplb.num_redundant_experts_per_rank logger.debug( f"global_rank {self.global_rank_} layerindex {self.layer_num_} " - f"redundancy_expertids: {self.redundancy_expert_ids}" - ) - self.local_n_routed_experts = self.n_routed_experts // self.global_world_size + self.redundancy_expert_num - n_experts_per_rank = self.n_routed_experts // self.global_world_size - start_expert_id = self.global_rank_ * n_experts_per_rank - self.local_expert_ids = ( - list(range(start_expert_id, start_expert_id + n_experts_per_rank)) + self.redundancy_expert_ids + f"initial_redundant_expert_ids: {self._initial_redundant_expert_ids}" ) + num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank + self.local_n_routed_experts = num_primary_experts_per_rank + num_redundant_experts_per_rank + start_expert_id = self.global_rank_ * num_primary_experts_per_rank self.expert_idx_to_local_idx = { - expert_idx: expert_idx - start_expert_id for expert_idx in self.local_expert_ids[:n_experts_per_rank] + expert_idx: expert_idx - start_expert_id + for expert_idx in range(start_expert_id, start_expert_id + num_primary_experts_per_rank) } - self.redundancy_expert_idx_to_local_idx = { - redundancy_expert_idx: n_experts_per_rank + i - for (i, redundancy_expert_idx) in enumerate(self.redundancy_expert_ids) + self._initial_redundant_expert_idx_to_local_idx = { + redundant_expert_idx: num_primary_experts_per_rank + i + for (i, redundant_expert_idx) in enumerate(self._initial_redundant_expert_ids) } else: self.local_expert_ids = list(range(self.n_routed_experts + self.num_fused_shared_experts)) self.expert_idx_to_local_idx = {expert_idx: i for (i, expert_idx) in enumerate(self.local_expert_ids)} - self.rexpert_idx_to_local_idx = {} def experts( self, @@ -172,7 +211,7 @@ def experts_with_topk( moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) if moe_capture_callback is not None: moe_capture_callback(topk_ids) - return self.fuse_moe_impl.fused_experts_with_topk( + return self.fuse_moe_impl._fused_experts( input_tensor=input_tensor, w13=self.w13, w2=self.w2, @@ -331,8 +370,8 @@ def load_hf_weights(self, weights): self._load_e_score_correction_bias(weights) self._load_per_expert_scale(weights) self._load_weight(self.expert_idx_to_local_idx, weights) - if self.redundancy_expert_num > 0: - self._load_weight(self.redundancy_expert_idx_to_local_idx, weights) + if self._initial_redundant_expert_idx_to_local_idx: + self._load_weight(self._initial_redundant_expert_idx_to_local_idx, weights) def verify_load(self): weight_load_ok = all(all(_weight_pack.load_ok) for _weight_pack in self.w1_list + self.w2_list + self.w3_list) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 9b32284a1a..89fc529801 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -3,18 +3,33 @@ from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM from .mxfp4_impl import FuseMoeMXFP4 +from ..expert_parallel_state import ExpertParallelState -def select_fuse_moe_impl(quant_method: QuantizationMethod, enable_ep_moe: bool): +def create_fuse_moe_impl( + *, + n_routed_experts: int, + num_fused_shared_experts: int, + routed_scaling_factor: float, + quant_method: QuantizationMethod, + expert_parallel_state: ExpertParallelState | None = None, +): if quant_method.method_name == "mxfp4w4a16-b32-marlin": - if enable_ep_moe: + if expert_parallel_state is not None: raise RuntimeError("mxfp4w4a16-b32-marlin does not support enable_ep_moe yet") - return FuseMoeMXFP4 - - if enable_ep_moe: - return FuseMoeDeepGEMM - - if quant_method.method_name == "awq_marlin": - return FuseMoeMarlin + impl_cls = FuseMoeMXFP4 + elif expert_parallel_state is not None: + impl_cls = FuseMoeDeepGEMM + elif quant_method.method_name == "awq_marlin": + impl_cls = FuseMoeMarlin else: - return FuseMoeTriton + impl_cls = FuseMoeTriton + kwargs = dict( + n_routed_experts=n_routed_experts, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + quant_method=quant_method, + ) + if expert_parallel_state is not None: + kwargs["expert_parallel_state"] = expert_parallel_state + return impl_cls(**kwargs) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index 1e3ad4b196..35e872df10 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -1,53 +1,25 @@ import torch -from abc import abstractmethod -from typing import Callable, Optional +from abc import ABC, abstractmethod +from typing import Callable, Optional, Tuple from lightllm.common.quantization.quantize_method import ( WeightPack, QuantizationMethod, ) -from lightllm.utils.dist_utils import ( - get_global_rank, - get_global_world_size, -) -class FuseMoeBaseImpl: +class FuseMoeBaseImpl(ABC): def __init__( self, n_routed_experts: int, num_fused_shared_experts: int, routed_scaling_factor: float, quant_method: QuantizationMethod, - redundancy_expert_num: int, - redundancy_expert_ids_tensor: torch.Tensor, - routed_expert_counter_tensor: torch.Tensor, - auto_update_redundancy_expert: bool, ): self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self.routed_scaling_factor = routed_scaling_factor self.quant_method = quant_method - self.global_rank_ = get_global_rank() - self.global_world_size_ = get_global_world_size() - self.ep_n_routed_experts = self.n_routed_experts // self.global_world_size_ - self.total_expert_num_contain_redundancy = ( - self.n_routed_experts + redundancy_expert_num * self.global_world_size_ - ) - - # redundancy expert related - self.redundancy_expert_num = redundancy_expert_num - self.redundancy_expert_ids_tensor = redundancy_expert_ids_tensor - self.routed_expert_counter_tensor = routed_expert_counter_tensor - self.auto_update_redundancy_expert = auto_update_redundancy_expert - # workspace for kernel optimization - self.workspace = self.create_workspace() - - @abstractmethod - def create_workspace(self): - pass - - @abstractmethod def __call__( self, input_tensor: torch.Tensor, @@ -67,5 +39,62 @@ def __call__( per_expert_scale: Optional[torch.Tensor] = None, # Qwen3.5 uses this gate to control fused shared expert aggregation weights. shared_expert_gate: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + topk_weights, topk_ids, origin_topk_ids = self._select_experts( + input_tensor=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + top_k=top_k, + renormalize=renormalize, + use_grouped_topk=use_grouped_topk, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + per_expert_scale=per_expert_scale, + shared_expert_gate=shared_expert_gate, + is_prefill=is_prefill, + preserve_logical_ids=moe_capture_callback is not None, + ) + if moe_capture_callback is not None: + moe_capture_callback(origin_topk_ids) + return self._fused_experts( + input_tensor=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + router_logits=router_logits, + is_prefill=is_prefill, + ) + + @abstractmethod + def _select_experts( + self, + input_tensor: torch.Tensor, + router_logits: torch.Tensor, + correction_bias: Optional[torch.Tensor], + top_k: int, + renormalize: bool, + use_grouped_topk: bool, + topk_group: int, + num_expert_group: int, + scoring_func: str, + per_expert_scale: Optional[torch.Tensor] = None, + shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + pass + + @abstractmethod + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, ) -> torch.Tensor: pass diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 417910f89c..627c09245d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -1,6 +1,7 @@ import torch from typing import Optional, Tuple, Any -from .triton_impl import FuseMoeTriton +from .base_impl import FuseMoeBaseImpl +from ..expert_parallel_state import ExpertParallelState from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( @@ -16,88 +17,15 @@ ) from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.common.triton_utils.autotuner import Autotuner -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair -class FuseMoeDeepGEMM(FuseMoeTriton): - def __init__(self, *args, **kwargs): +class FuseMoeDeepGEMM(FuseMoeBaseImpl): + def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): super().__init__(*args, **kwargs) + self.expert_parallel_state = expert_parallel_state + self.eplb = expert_parallel_state.eplb self.ep_balance_counters = None - - def _select_experts( - self, - input_tensor: torch.Tensor, - router_logits: torch.Tensor, - correction_bias: Optional[torch.Tensor], - top_k: int, - renormalize: bool, - use_grouped_topk: bool, - topk_group: int, - num_expert_group: int, - scoring_func: str, - per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - ): - """Select experts and return topk weights and ids.""" - assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts - - topk_weights, topk_ids = select_experts( - hidden_states=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - use_grouped_topk=use_grouped_topk, - top_k=top_k, - renormalize=renormalize, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - ) - if self.routed_scaling_factor != 1.0: - topk_weights.mul_(self.routed_scaling_factor) - if per_expert_scale is not None: - topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) - origin_topk_ids = topk_ids - if self.redundancy_expert_num > 0: - # 因为 redundancy_topk_ids_repair 会修改 topk_ids,所以需要先复制一份 - origin_topk_ids = topk_ids.clone() - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=self.redundancy_expert_ids_tensor, - ep_expert_num=self.ep_n_routed_experts, - global_rank=self.global_rank_, - expert_counter=self.routed_expert_counter_tensor, - enable_counter=self.auto_update_redundancy_expert, - ) - return topk_weights, topk_ids, origin_topk_ids - - def _fused_experts( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, - alloc_tensor_func=torch.empty, - ): - output = fused_experts( - hidden_states=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_idx=topk_ids.to(torch.long), - num_experts=self.total_expert_num_contain_redundancy, # number of all experts contain redundancy - quant_method=self.quant_method, - is_prefill=is_prefill, - previous_event=None, # for overlap - clamp_limit=clamp_limit, - alloc_tensor_func=alloc_tensor_func, - ep_balance_counters=self.ep_balance_counters, - ) - return output + self._primary_weight_pack_cache = {} def low_latency_dispatch( self, @@ -121,6 +49,7 @@ def low_latency_dispatch( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, + is_prefill=False, ) return self.low_latency_dispatch_with_topk( @@ -142,7 +71,8 @@ def low_latency_dispatch_with_topk( topk_idx=topk_idx, x=hidden_states, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - num_experts=self.total_expert_num_contain_redundancy, + # decode 与 EPLB 的物理冗余行刻意隔离:DeepEP 使用原始 logical expert ID。 + num_experts=self.n_routed_experts, use_fp8=use_fp8_w8a8, async_finish=False, return_recv_hook=True, @@ -175,6 +105,7 @@ def select_experts_and_quant_input( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, + is_prefill=True, ) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor @@ -192,7 +123,7 @@ def dispatch( qinput_tensor, topk_idx=topk_idx, topk_weights=topk_weights, - num_experts=self.total_expert_num_contain_redundancy, + num_experts=self.expert_parallel_state.num_total_physical_experts, num_max_tokens_per_rank=num_max_tokens_per_rank, expert_alignment=128, num_sms=get_ep_num_sms(), @@ -235,6 +166,7 @@ def masked_group_gemm( expected_m: int, clamp_limit: Optional[float] = None, ): + w13, w2 = self._primary_weight_pack(w13), self._primary_weight_pack(w2) w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale return masked_group_gemm( @@ -335,3 +267,124 @@ def hook(): event.current_stream_wait() return combined_x, hook + + def _select_experts( + self, + input_tensor: torch.Tensor, + router_logits: torch.Tensor, + correction_bias: Optional[torch.Tensor], + top_k: int, + renormalize: bool, + use_grouped_topk: bool, + topk_group: int, + num_expert_group: int, + scoring_func: str, + per_expert_scale: Optional[torch.Tensor] = None, + shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, + ): + """选择 expert;EPLB prefill 统一由融合路径返回 physical ID。""" + assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" + eplb = self.eplb + if is_prefill is True and eplb is not None: + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb + + group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 + topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( + hidden_states=input_tensor, + gating_output=router_logits, + correction_bias=correction_bias, + topk=top_k, + renormalize=renormalize, + num_expert_group=num_expert_group, + topk_group=topk_group, + scoring_func=scoring_func, + logical_to_physical_map=eplb.logical_to_physical_map, + logical_replica_count=eplb.logical_replica_count, + expert_counter=eplb.route_counter, + sample_index=eplb.next_sample_index(), + record_load=eplb.recording, + use_grouped_topk=use_grouped_topk, + return_logical_ids=preserve_logical_ids, + group_score_used_topk_num=group_score_topk_num, + ) + origin_topk_ids = logical_topk_ids if logical_topk_ids is not None else topk_ids + else: + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + + topk_weights, topk_ids = select_experts( + hidden_states=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + use_grouped_topk=use_grouped_topk, + top_k=top_k, + renormalize=renormalize, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + ) + origin_topk_ids = topk_ids + if self.routed_scaling_factor != 1.0: + topk_weights.mul_(self.routed_scaling_factor) + return topk_weights, topk_ids, origin_topk_ids + + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, + ): + if is_prefill is False: + w13 = self._primary_weight_pack(w13) + w2 = self._primary_weight_pack(w2) + num_experts = self.n_routed_experts + else: + num_experts = self.expert_parallel_state.num_total_physical_experts + + output = fused_experts( + hidden_states=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_idx=topk_ids.to(torch.long), + num_experts=num_experts, + quant_method=self.quant_method, + is_prefill=is_prefill, + previous_event=None, # for overlap + clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, + ep_balance_counters=self.ep_balance_counters, + ) + return output + + def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: + """返回所有 decode 路径使用的缓存本地主副本视图。""" + if self.eplb is None: + return weight_pack + cache = self._primary_weight_pack_cache + cache_key = id(weight_pack) + primary = cache.get(cache_key) + if primary is None: + num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank + primary = WeightPack( + weight=weight_pack.weight[:num_primary_experts_per_rank], + weight_scale=( + weight_pack.weight_scale[:num_primary_experts_per_rank] + if weight_pack.weight_scale is not None + else None + ), + weight_zero_point=( + getattr(weight_pack, "weight_zero_point", None)[:num_primary_experts_per_rank] + if getattr(weight_pack, "weight_zero_point", None) is not None + else None + ), + ) + cache[cache_key] = primary + return primary diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index a30a669c18..59ba9fe766 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -11,6 +11,10 @@ class FuseMoeMarlin(FuseMoeTriton): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.workspace = self.create_workspace() + def create_workspace(self): from lightllm.utils.vllm_utils import HAS_VLLM diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index a44bf16c54..43ea96865e 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -1,36 +1,10 @@ import torch -from typing import Callable, Optional +from typing import Optional from lightllm.common.quantization.no_quant import WeightPack -from lightllm.common.quantization.quantize_method import QuantizationMethod from .base_impl import FuseMoeBaseImpl class FuseMoeTriton(FuseMoeBaseImpl): - def __init__( - self, - n_routed_experts: int, - num_fused_shared_experts: int, - routed_scaling_factor: float, - quant_method: QuantizationMethod, - redundancy_expert_num: int, - redundancy_expert_ids_tensor: torch.Tensor, - routed_expert_counter_tensor: torch.Tensor, - auto_update_redundancy_expert: bool, - ): - super().__init__( - n_routed_experts=n_routed_experts, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - quant_method=quant_method, - redundancy_expert_num=redundancy_expert_num, - redundancy_expert_ids_tensor=redundancy_expert_ids_tensor, - routed_expert_counter_tensor=routed_expert_counter_tensor, - auto_update_redundancy_expert=auto_update_redundancy_expert, - ) - - def create_workspace(self): - return None - def _select_experts( self, input_tensor: torch.Tensor, @@ -44,6 +18,8 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + preserve_logical_ids: bool = False, ): """Select experts and return topk weights and ids.""" from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts @@ -109,72 +85,3 @@ def _fused_experts( limit=clamp_limit, ) return input_tensor - - def fused_experts_with_topk( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, - alloc_tensor_func=torch.empty, - ): - return self._fused_experts( - input_tensor=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_ids=topk_ids, - is_prefill=is_prefill, - clamp_limit=clamp_limit, - alloc_tensor_func=alloc_tensor_func, - ) - - def __call__( - self, - input_tensor: torch.Tensor, - router_logits: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - correction_bias: Optional[torch.Tensor], - scoring_func: str, - top_k: int, - renormalize: bool, - use_grouped_topk: bool, - topk_group: int, - num_expert_group: int, - is_prefill: Optional[bool] = None, - # Callback to capture MoE topk expert ids (routed experts metadata). - moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None, - per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - ): - topk_weights, topk_ids, origin_topk_ids = self._select_experts( - input_tensor=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - top_k=top_k, - renormalize=renormalize, - use_grouped_topk=use_grouped_topk, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - per_expert_scale=per_expert_scale, - shared_expert_gate=shared_expert_gate, - ) - - if moe_capture_callback is not None: - moe_capture_callback(origin_topk_ids) - - output = self._fused_experts( - input_tensor=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_ids=topk_ids, - router_logits=router_logits, - is_prefill=is_prefill, - ) - return output diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index fb0323cd4b..8544b22e32 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -5,6 +5,14 @@ from triton.language.standard import _log2, sum, zeros_like +@triton.jit +def _eplb_replica_index(token_index, logical_id, replica_count): + """Choose a replica with independent phases for a token's top-k experts.""" + token_hash = token_index.to(tl.uint32) * 2654435769 + expert_hash = logical_id.to(tl.uint32) * 2246822519 + return (token_hash + expert_hash) % replica_count.to(tl.uint32) + + @triton.jit def _compare_and_swap(x, x_1, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): n_outer: tl.core.constexpr = x.numel >> n_dims @@ -202,6 +210,172 @@ def grouped_topk_kernel( return +@triton.jit +def grouped_topk_eplb_kernel( + gating_output_ptr, + gating_output_stride_m, + gating_output_stride_n, + correction_bias_ptr, + out_topk_weights, + out_topk_weights_stride_m, + out_topk_weights_stride_n, + out_topk_ids, + out_topk_ids_stride_m, + out_topk_ids_stride_n, + out_logical_ids, + out_logical_ids_stride_m, + out_logical_ids_stride_n, + logical_to_physical_ptr, + logical_replica_count_ptr, + expert_counter_ptr, + sample_index, + group_num, + group_expert_num, + total_expert_num, + group_topk_num, + IS_SIGMOID: tl.constexpr, + USE_GROUPED_TOPK: tl.constexpr, + HAS_CORRECTION_BIAS: tl.constexpr, + RETURN_LOGICAL_IDS: tl.constexpr, + EXPERT_GROUP_NUM: tl.constexpr, + EXPERT_GROUP_SIZE: tl.constexpr, + TOPK_NUM: tl.constexpr, + TOPK_BLOCK_SIZE: tl.constexpr, + RENORMALIZE: tl.constexpr, + GROUP_SCORE_USED_TOPK_NUM: tl.constexpr, + COUNTER_NUM_EXPERTS: tl.constexpr, + MAP_SLOTS: tl.constexpr, + RECORD_LOAD: tl.constexpr, + SINGLE_TOKEN: tl.constexpr, +): + """Grouped top-k, EPLB accounting, and replica mapping without a global score workspace.""" + token_index = tl.program_id(axis=0) + offs_group = tl.arange(0, EXPERT_GROUP_NUM) + offs_group_v = tl.arange(0, EXPERT_GROUP_SIZE) + logical_ids = offs_group[:, None] * group_expert_num + offs_group_v[None, :] + valid_expert = ( + (offs_group < group_num)[:, None] + & (offs_group_v < group_expert_num)[None, :] + & (logical_ids < total_expert_num) + ) + hidden_states = tl.load( + gating_output_ptr + token_index * gating_output_stride_m + logical_ids * gating_output_stride_n, + mask=valid_expert, + other=-float("inf"), + ).to(tl.float32) + + if IS_SIGMOID: + old_scores = tl.sigmoid(hidden_states) + else: + group_max = tl.max(hidden_states, axis=1) + global_max = tl.max(group_max, axis=0) + numerators = tl.where(valid_expert, tl.exp(hidden_states - global_max), 0.0) + denominator = tl.sum(tl.sum(numerators, axis=1), axis=0) + old_scores = numerators / denominator + + if HAS_CORRECTION_BIAS: + correction_bias = tl.load(correction_bias_ptr + logical_ids, mask=valid_expert, other=0.0) + scores = tl.where(valid_expert, old_scores + correction_bias, -float("inf")) + else: + scores = tl.where(valid_expert, old_scores, -float("inf")) + + if USE_GROUPED_TOPK: + if GROUP_SCORE_USED_TOPK_NUM == 1: + group_value = tl.max(scores, axis=1) + elif GROUP_SCORE_USED_TOPK_NUM == 2: + first_score, first_index = tl.max(scores, axis=1, return_indices=True) + second_score = tl.max( + tl.where(offs_group_v[None, :] == first_index[:, None], -float("inf"), scores), + axis=1, + ) + group_value = first_score + second_score + else: + sorted_group_scores = tl.sort(scores, dim=1, descending=True) + group_value = tl.sum( + tl.where(offs_group_v[None, :] < GROUP_SCORE_USED_TOPK_NUM, sorted_group_scores, 0.0), + axis=1, + ) + + if EXPERT_GROUP_NUM > 1: + sorted_group_value = tl.sort(group_value, descending=True) + else: + sorted_group_value = group_value + group_topk_value = tl.sum(tl.where(offs_group == group_topk_num - 1, sorted_group_value, 0.0)) + candidate_scores = tl.where( + (group_value >= group_topk_value)[:, None] & valid_expert, + scores, + -float("inf"), + ) + else: + candidate_scores = tl.where(valid_expert, old_scores, -float("inf")) + + sort_block_size: tl.constexpr = EXPERT_GROUP_NUM * EXPERT_GROUP_SIZE + flat_offsets = tl.arange(0, sort_block_size) + candidate_scores = tl.reshape(candidate_scores, (sort_block_size,)) + topk_offsets = tl.arange(0, TOPK_BLOCK_SIZE) + selected_weights = tl.zeros((TOPK_BLOCK_SIZE,), tl.float32) + selected_logical_ids = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) + sum_scores = 0.0 + for topk_index in range(TOPK_NUM): + selected_offset = tl.argmax(candidate_scores, axis=0) + selected_group = selected_offset // EXPERT_GROUP_SIZE + selected_group_offset = selected_offset % EXPERT_GROUP_SIZE + selected_logical_id = selected_group * group_expert_num + selected_group_offset + selected_hidden_state = tl.load( + gating_output_ptr + token_index * gating_output_stride_m + selected_logical_id * gating_output_stride_n + ).to(tl.float32) + if IS_SIGMOID: + selected_weight = tl.sigmoid(selected_hidden_state) + else: + selected_weight = tl.exp(selected_hidden_state - global_max) / denominator + sum_scores += selected_weight + topk_lane = topk_offsets == topk_index + selected_weights = tl.where(topk_lane, selected_weight, selected_weights) + selected_logical_ids = tl.where(topk_lane, selected_logical_id, selected_logical_ids) + candidate_scores = tl.where(flat_offsets == selected_offset, -float("inf"), candidate_scores) + + topk_mask = topk_offsets < TOPK_NUM + if RECORD_LOAD: + tl.atomic_add( + expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + selected_logical_ids, + 1, + mask=topk_mask, + sem="relaxed", + ) + if SINGLE_TOKEN: + replica_indices = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) + else: + replica_counts = tl.load( + logical_replica_count_ptr + selected_logical_ids, + mask=topk_mask, + other=1, + ) + replica_indices = _eplb_replica_index(token_index, selected_logical_ids, replica_counts) + selected_physical_ids = tl.load( + logical_to_physical_ptr + selected_logical_ids * MAP_SLOTS + replica_indices, + mask=topk_mask, + other=-1, + ) + if RENORMALIZE: + selected_weights /= sum_scores + tl.store( + out_topk_weights + token_index * out_topk_weights_stride_m + topk_offsets * out_topk_weights_stride_n, + selected_weights, + mask=topk_mask, + ) + tl.store( + out_topk_ids + token_index * out_topk_ids_stride_m + topk_offsets * out_topk_ids_stride_n, + selected_physical_ids, + mask=topk_mask, + ) + if RETURN_LOGICAL_IDS: + tl.store( + out_logical_ids + token_index * out_logical_ids_stride_m + topk_offsets * out_logical_ids_stride_n, + selected_logical_ids, + mask=topk_mask, + ) + + def triton_grouped_topk( hidden_states: torch.Tensor, gating_output: torch.Tensor, @@ -263,3 +437,81 @@ def triton_grouped_topk( num_stages=1, ) return out_topk_weights, out_topk_ids + + +def triton_grouped_topk_eplb( + hidden_states: torch.Tensor, + gating_output: torch.Tensor, + correction_bias: torch.Tensor, + topk: int, + renormalize: bool, + num_expert_group: int, + topk_group: int, + scoring_func: str, + logical_to_physical_map: torch.Tensor, + logical_replica_count: torch.Tensor, + expert_counter: torch.Tensor, + sample_index: int, + record_load: bool, + use_grouped_topk: bool, + return_logical_ids: bool = False, + group_score_used_topk_num: int = 2, +): + """Fused EPLB prefill top-k returning physical IDs and optional logical IDs.""" + token_num, total_expert_num = gating_output.shape + out_topk_weights = torch.empty((token_num, topk), dtype=torch.float32, device=gating_output.device) + out_topk_ids = torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) + out_logical_ids = ( + torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) if return_logical_ids else None + ) + if token_num == 0: + return out_topk_weights, out_topk_ids, out_logical_ids + if use_grouped_topk: + assert total_expert_num % num_expert_group == 0 + group_num = num_expert_group + group_expert_num = total_expert_num // num_expert_group + group_topk_num = topk_group + else: + group_num = 1 + group_expert_num = total_expert_num + group_topk_num = 1 + expert_group_num = triton.next_power_of_2(group_num) + expert_group_size = triton.next_power_of_2(group_expert_num) + sort_block_size = expert_group_num * expert_group_size + num_warps = min(max(1, sort_block_size // 256), 8) + grouped_topk_eplb_kernel[(token_num,)]( + gating_output, + *gating_output.stride(), + correction_bias, + out_topk_weights, + *out_topk_weights.stride(), + out_topk_ids, + *out_topk_ids.stride(), + out_logical_ids if out_logical_ids is not None else out_topk_ids, + *(out_logical_ids.stride() if out_logical_ids is not None else out_topk_ids.stride()), + logical_to_physical_map, + logical_replica_count, + expert_counter, + sample_index, + group_num=group_num, + group_expert_num=group_expert_num, + total_expert_num=total_expert_num, + group_topk_num=group_topk_num, + IS_SIGMOID=use_grouped_topk and scoring_func == "sigmoid", + USE_GROUPED_TOPK=use_grouped_topk, + HAS_CORRECTION_BIAS=use_grouped_topk and correction_bias is not None, + RETURN_LOGICAL_IDS=return_logical_ids, + EXPERT_GROUP_NUM=expert_group_num, + EXPERT_GROUP_SIZE=expert_group_size, + TOPK_NUM=topk, + TOPK_BLOCK_SIZE=triton.next_power_of_2(topk), + RENORMALIZE=renormalize, + GROUP_SCORE_USED_TOPK_NUM=group_score_used_topk_num, + COUNTER_NUM_EXPERTS=expert_counter.shape[1], + MAP_SLOTS=logical_to_physical_map.shape[1], + RECORD_LOAD=record_load, + SINGLE_TOKEN=token_num == 1, + num_warps=num_warps, + num_stages=1, + ) + return out_topk_weights, out_topk_ids, out_logical_ids diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py index 1c01cbd638..d2f59de480 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py @@ -21,7 +21,6 @@ from lightllm.utils.sgl_utils import sgl_ops from typing import Callable, List, Optional, Tuple from lightllm.common.basemodel.triton_kernel.fused_moe.softmax_topk import softmax_topk -from lightllm.common.triton_utils.autotuner import Autotuner def fused_topk( @@ -168,12 +167,4 @@ def select_experts( hidden_states=hidden_states, gating_output=router_logits, topk=top_k, renormalize=renormalize ) - ######################################## warning ################################################## - # here is used to match autotune feature, make topk_ids more random - if Autotuner.is_autotune_warmup(): - rand_gen = torch.Generator(device="cuda") - rand_gen.manual_seed(router_logits.shape[0]) - router_logits = torch.randn(size=router_logits.shape, generator=rand_gen, dtype=torch.float32, device="cuda") - _, topk_ids = torch.topk(router_logits, k=top_k, dim=1) - return topk_weights, topk_ids diff --git a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py deleted file mode 100644 index ba48f414db..0000000000 --- a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py +++ /dev/null @@ -1,111 +0,0 @@ -import torch -import triton -import triton.language as tl - - -@triton.jit -def _redundancy_topk_ids_repair_kernel( - topk_ids_ptr, - topk_total_num, - ep_expert_num, - redundancy_expert_num, - global_rank, - redundancy_expert_ids_ptr, - expert_counter_ptr, - BLOCK_SIZE: tl.constexpr, - ENABLE_COUNTER: tl.constexpr, -): - block_index = tl.program_id(0) - offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offs_d < topk_total_num - current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) - - if ENABLE_COUNTER: - tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) - - # Remap original expert IDs to a new space that accounts for redundant expert slots. - new_current_topk_ids = (current_topk_ids // ep_expert_num) * redundancy_expert_num + current_topk_ids - - for i in tl.range(0, redundancy_expert_num, step=1, num_stages=3): - cur_redundancy_expert_id = tl.load(redundancy_expert_ids_ptr + i) - cur_redundancy_expert_id = ( - cur_redundancy_expert_id // ep_expert_num - ) * redundancy_expert_num + cur_redundancy_expert_id - new_current_topk_ids = tl.where( - new_current_topk_ids == cur_redundancy_expert_id, - (ep_expert_num + redundancy_expert_num) * (global_rank) + ep_expert_num + i, - new_current_topk_ids, - ) - - tl.store(topk_ids_ptr + offs_d, new_current_topk_ids, mask=mask) - return - - -@torch.no_grad() -def redundancy_topk_ids_repair( - topk_ids: torch.Tensor, - redundancy_expert_ids: torch.Tensor, - ep_expert_num: int, - global_rank: int, - expert_counter: torch.Tensor = None, - enable_counter: bool = False, -): - assert topk_ids.is_contiguous() - assert len(topk_ids.shape) == 2 - assert redundancy_expert_ids is not None - redundancy_expert_num = redundancy_expert_ids.shape[0] - BLOCK_SIZE = 512 - grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) - num_warps = 4 - - _redundancy_topk_ids_repair_kernel[grid]( - topk_ids_ptr=topk_ids, - topk_total_num=topk_ids.numel(), - ep_expert_num=ep_expert_num, - redundancy_expert_num=redundancy_expert_num, - global_rank=global_rank, - redundancy_expert_ids_ptr=redundancy_expert_ids, - expert_counter_ptr=expert_counter, - BLOCK_SIZE=BLOCK_SIZE, - ENABLE_COUNTER=enable_counter, - num_warps=num_warps, - num_stages=3, - ) - return - - -@triton.jit -def _expert_id_counter_kernel( - topk_ids_ptr, - topk_total_num, - expert_counter_ptr, - BLOCK_SIZE: tl.constexpr, -): - block_index = tl.program_id(0) - offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offs_d < topk_total_num - current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) - tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) - return - - -@torch.no_grad() -def expert_id_counter( - topk_ids: torch.Tensor, - expert_counter: torch.Tensor, -): - assert topk_ids.is_contiguous() - assert len(topk_ids.shape) == 2 - BLOCK_SIZE = 512 - grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) - num_warps = 4 - - _expert_id_counter_kernel[grid]( - topk_ids_ptr=topk_ids, - topk_total_num=topk_ids.numel(), - expert_counter_ptr=expert_counter, - BLOCK_SIZE=BLOCK_SIZE, - num_warps=num_warps, - num_stages=1, - ) - return diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py new file mode 100644 index 0000000000..441f92ca3e --- /dev/null +++ b/lightllm/common/eplb_utils.py @@ -0,0 +1,16 @@ +"""Small, dependency-free EPLB helpers shared by transfer and model profiling.""" + + +EPLB_MAX_STAGING_DEPTH = 8 + + +def extract_eplb_expert_tensors(weight): + result = [] + for pack_name in ("w13", "w2"): + pack = getattr(weight, pack_name) + for value_name in ("weight", "weight_scale", "weight_zero_point"): + tensor = getattr(pack, value_name, None) + if tensor is not None: + assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" + result.append((f"{pack_name}.{value_name}", tensor)) + return result diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index f473e031ad..b2f13ecc04 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -26,13 +26,30 @@ class MemoryManager: operator_class = NormalMemOperator - def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): + def __init__( + self, + size, + dtype, + head_num, + head_dim, + layer_num, + always_copy=False, + mem_fraction=0.9, + memory_reservations=None, + ): self.size = size self.head_num = head_num self.head_dim = head_dim self.layer_num = layer_num self.always_copy = always_copy self.dtype = dtype + # Named reservations are allocations made after KV profiling. They are + # deliberately outside get_fixed_memory_size(): model-specific exact KV + # geometry owns fixed bytes, while these values are deducted once from + # the profile budget only. + self.memory_reservations = dict(memory_reservations or {}) + if any(value < 0 for value in self.memory_reservations.values()): + raise ValueError(f"memory reservations must be non-negative: {self.memory_reservations}") # profile the max total token num if the size is None self.profile_size(mem_fraction) @@ -70,11 +87,14 @@ def profile_size(self, mem_fraction): available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) cell_size = self.get_cell_size() fixed_memory_size = self.get_fixed_memory_size() - available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size + reservations = getattr(self, "memory_reservations", {}) + reserved_memory_size = sum(reservations.values()) + available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size if available_memory_bytes <= 0: raise RuntimeError( f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " - f"but only {available_memory:.2f} GB is available for KV cache" + f"plus {reserved_memory_size / 1024**3:.2f} GB reservations, " + f"but only {available_memory:.2f} GB is available" ) self.size = int(available_memory_bytes / cell_size) if world_size > 1: @@ -84,6 +104,7 @@ def profile_size(self, mem_fraction): logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" + f"{reservations} bytes are reserved for post-profile model buffers\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 08015789a6..5b9e64ac58 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -29,7 +29,6 @@ get_env_start_args, get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, - get_redundancy_expert_num, ) from lightllm.utils.dist_utils import ( get_global_world_size, @@ -272,8 +271,15 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - redundancy_expert_num = get_redundancy_expert_num() - self.ll_num_experts = n_routed_experts + redundancy_expert_num * global_world_size + total_redundant_experts = ( + get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size + if get_env_start_args().enable_prefill_eplb + else 0 + ) + self.ll_prefill_num_experts = n_routed_experts + total_redundant_experts + # EPLB's redundant rows are a prefill-only physical layout; decode + # always routes the logical expert space. + self.ll_decode_num_experts = n_routed_experts self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, @@ -298,7 +304,7 @@ def new_deepep_group( enable_env_vars("LIGHTLLM_ENABLE_SM90_FP8_MEGA_MOE") and is_sm90_gpu() and FP8_MOE_QUANT_METHOD in expert_quant_method_names - and redundancy_expert_num == 0 + and total_redundant_experts == 0 ): self.ep_mega_moe_mma_type = "fp8xfp8" self.ep_mega_moe_quant_method = FP8_MOE_QUANT_METHOD @@ -336,7 +342,10 @@ def new_deepep_group( # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( - self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts + self.ll_decode_num_tokens, + self.ll_hidden, + global_world_size, + self.ll_decode_num_experts, ) num_rdma_bytes = decode_size_hint # normal 节点同时执行 Prefill 和 Decode,复用的 RDMA buffer 必须覆盖全部 Prefill workspace。 @@ -346,7 +355,7 @@ def new_deepep_group( hidden_size=self.ll_hidden, intermediate_size=moe_intermediate_size, num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_num_experts, + num_experts=self.ll_prefill_num_experts, world_size=global_world_size, hidden_dtype=get_torch_dtype(args.data_type), ) @@ -355,7 +364,7 @@ def new_deepep_group( deepep_group, num_rdma_bytes=num_rdma_bytes, low_latency_mode=True, - num_qps_per_rank=(self.ll_num_experts // global_world_size), + num_qps_per_rank=(self.ll_decode_num_experts // global_world_size), ) self.ep_prefill_moe_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( torch.uint8, use_rdma_buffer=True @@ -367,7 +376,7 @@ def new_deepep_group( hidden_size=self.ll_hidden, intermediate_size=moe_intermediate_size, num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_num_experts, + num_experts=self.ll_prefill_num_experts, world_size=global_world_size, hidden_dtype=get_torch_dtype(args.data_type), ) @@ -377,9 +386,10 @@ def new_deepep_group( device=torch.device("cuda", torch.cuda.current_device()), ) - theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_num_experts, num_experts_per_tok) + theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_prefill_num_experts, num_experts_per_tok) deepep_sms = 0 if self.ep_mega_moe_mma_type == "fp8xfp8" and not has_legacy_moe_layer else theoretical_sms - self._set_num_sms_for_deep_gemm(deepep_sms) + low_latency_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_decode_num_experts, num_experts_per_tok) + self._set_num_sms_for_deep_gemm(deepep_sms, low_latency_sms) if enable_mega_moe_buffer: if moe_intermediate_size is None: @@ -392,7 +402,7 @@ def new_deepep_group( ) self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, - self.ll_num_experts, + self.ll_decode_num_experts, self.ll_num_tokens, num_experts_per_tok, self.ll_hidden, @@ -401,15 +411,18 @@ def new_deepep_group( ) logger.info( "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " - "mega_moe=%s, mega_moe_mma_type=%s, expert_quant_method_names=%s", + "mega_moe=%s, mega_moe_mma_type=%s, ll_prefill_num_experts=%s, " + "ll_decode_num_experts=%s, expert_quant_method_names=%s", enable_low_latency_buffer, self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, self.ep_mega_moe_mma_type, + self.ll_prefill_num_experts, + self.ll_decode_num_experts, sorted(expert_quant_method_names), ) - def _set_num_sms_for_deep_gemm(self, deepep_sms: int): + def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): try: try: from deep_gemm.jit_kernels.utils import set_num_sms @@ -418,9 +431,12 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): device_sms = get_device_sm_count() deepep_sms = max(0, min(deepep_sms, max(device_sms - 2, 0))) + low_latency_sms = max(0, min(low_latency_sms, max(device_sms - 2, 0))) self.ep_num_sms = deepep_sms if self.ep_low_latency_buffer is not None: - deep_ep.Buffer.set_num_sms(deepep_sms - deepep_sms % 2) + # This setting controls the legacy low-latency buffer; keep + # its SM reservation based on decode's logical expert count. + deep_ep.Buffer.set_num_sms(low_latency_sms - low_latency_sms % 2) deep_gemm_sms = max(device_sms - deepep_sms, 2) if self.ep_mega_moe_mma_type == "fp8xfp8": deep_gemm_sms -= deep_gemm_sms % 2 @@ -467,7 +483,7 @@ def clear_deepep_buffer(self): """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( - self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts + self.ll_decode_num_tokens, self.ll_hidden, self.ll_decode_num_experts ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index f0eef53a4d..f43d8fe8b7 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -755,15 +755,16 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", ) parser.add_argument( - "--ep_redundancy_expert_config_path", - type=str, - default=None, - help="""Path of the redundant expert config. It can be used for deepseekv3 model.""", + "--enable_prefill_eplb", + action="store_true", + help="""Enable online expert load balancing for prefill only.""", ) parser.add_argument( - "--auto_update_redundancy_expert", - action="store_true", - help="""Whether to update the redundant expert for deepseekv3 model by online expert used counter.""", + "--eplb_num_redundant_experts_per_rank", + type=int, + default=2, + help="""Number of redundant physical experts per EP rank for each MoE layer used by prefill EPLB. + The value must be greater than 0.""", ) parser.add_argument( "--enable_fused_shared_experts", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index cf0154397d..06e15ed6f4 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -27,6 +27,7 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args +from lightllm.utils.device_utils import is_sm100_gpu from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup logger = init_logger(__name__) @@ -178,6 +179,15 @@ def _launch_subprocesses(args: StartArgs): if args.enable_dp_prefill_balance: assert args.enable_tpsp_mix_mode and args.dp > 1, "need set --enable_tpsp_mix_mode firstly and --dp > 1" + if args.enable_prefill_eplb: + assert args.enable_ep_moe, "--enable_prefill_eplb requires --enable_ep_moe" + assert not args.enable_prefill_cudagraph, "--enable_prefill_eplb does not support --enable_prefill_cudagraph" + # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. + assert not is_sm100_gpu(), "--enable_prefill_eplb does not support SM100" + assert ( + args.eplb_num_redundant_experts_per_rank > 0 + ), "--eplb_num_redundant_experts_per_rank must be greater than 0" + if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} for backend in args.llm_prefill_att_backend: diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 4d5e6b175e..3d036f55ad 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -183,8 +183,8 @@ class StartArgs: ) enable_ep_moe: bool = field(default=False) disable_ep_balance_monitor: bool = field(default=False) - ep_redundancy_expert_config_path: Optional[str] = field(default=None) - auto_update_redundancy_expert: bool = field(default=False) + enable_prefill_eplb: bool = field(default=False) + eplb_num_redundant_experts_per_rank: int = field(default=2) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, 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 5b3cd12c39..43032b15ae 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -13,6 +13,9 @@ 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 +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + disable_eplb_model_init, +) from lightllm.common.req_manager import DeepseekV4ReqManager, ReqManagerForMamba from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager @@ -274,6 +277,11 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None + if self.args.enable_prefill_eplb: + from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager + + self.model.eplb_manager = EPLBManager(self.model) + dist.barrier() # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 @@ -362,7 +370,9 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_cfg=draft_model_cfg, spec_mode=spec_mode, ) - self.draft_models.append(draft_model_class(draft_model_kvargs)) + with disable_eplb_model_init(): + draft_model = draft_model_class(draft_model_kvargs) + self.draft_models.append(draft_model) self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index 875bd5b738..222983a892 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 @@ -61,6 +61,11 @@ def infer_loop(self): event_pack.wait_to_forward() + # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal + # collectives/forward. + if self.model.eplb_manager is not None: + self.model.eplb_manager.poll() + self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( 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 c712d87ecb..b6e3a4b917 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 @@ -128,6 +128,11 @@ def infer_loop(self): event_pack.wait_to_forward() + # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal + # collectives/forward. + if self.model.eplb_manager is not None: + self.model.eplb_manager.poll() + self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py new file mode 100644 index 0000000000..a73f9a54b7 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -0,0 +1,550 @@ +import threading +import time +from typing import Dict, Optional + +import torch +import torch.distributed as dist + +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_logical_to_physical_maps_for_layers, + plan_redundant_experts, + select_improving_placements, +) +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + NixlEPLBTransfer, + align_target_placement, + build_transfer_plan, +) +from lightllm.utils.dist_utils import get_global_rank, get_global_world_size, get_node_world_size +from lightllm.utils.envs_utils import ( + get_eplb_placement_stickiness, + get_eplb_rebalance_gain_threshold, + get_prefill_eplb_step_interval, +) +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) +EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 +EPLB_EXPERT_ALIGNMENT = 128 +EPLB_CONTROL_ERROR = -1 +EPLB_STEADY_SAMPLE_STEPS = 4 + + +class EPLBManager: + """Online EPLB with asynchronous GPU expert migration.""" + + def __init__(self, model: TpPartBaseModel): + self.weights = _find_fused_moe_weights(model) + assert self.weights, "EPLB requires at least one EP MoE layer" + self.global_rank = get_global_rank() + self.world_size = get_global_world_size() + self.node_world_size = get_node_world_size() + self._eplb_states = [weight.expert_parallel_state.eplb for weight in self.weights] + self.step_interval = get_prefill_eplb_step_interval() + self.rebalance_gain_threshold = get_eplb_rebalance_gain_threshold() + self.placement_stickiness = get_eplb_placement_stickiness() + self.sampling_interval = self.step_interval + self.prefill_steps = 0 + routed = {weight.expert_parallel_state.num_logical_experts for weight in self.weights} + redundant = {state.num_redundant_experts_per_rank for state in self._eplb_states} + assert len(routed) == len(redundant) == 1 + self.num_logical_experts = routed.pop() + self.num_redundant_experts_per_rank = redundant.pop() + self.current_placement = torch.stack( + [state.initial_redundant_expert_ids_by_rank for state in self._eplb_states] + ) + self.in_flight = False + self.target_placement = None + self.target_metadata = None + self.in_flight_started_at = None + self.evaluation_in_flight = False + self._evaluation_lock = threading.Lock() + self._evaluation_result = None + self._evaluation_error = None + self._evaluation_thread = None + # A fresh manager starts with one continuous base window. After a + # sufficient evaluation, steady state returns to the cheap sparse + # probe. An insufficient sparse probe schedules one fresh continuous + # base window before the next fixed sampling boundary. + self._continuous_collection_start_step: Optional[int] = None + self._continuous_collection_end_step: Optional[int] = self.step_interval + self._steady_collection_end_step: Optional[int] = None + self._reset_recorded_samples() + self._set_recording(True) + # Keep background evaluation collectives separate from the main-thread + # control/poll collectives: their ordering is intentionally independent. + self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") + self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") + self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") + # This control-group scalar is only touched from the main inference + # thread, never by the background evaluation thread. + self._control_ready_count = torch.empty(1, dtype=torch.int32) + self.transfer = NixlEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) + if self.global_rank == 0: + logger.info( + "eplb enabled " + f"layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " + f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " + f"step_interval={self.step_interval} " + f"rebalance_gain_threshold={self.rebalance_gain_threshold:.4f} " + f"placement_stickiness={self.placement_stickiness:.4f}" + ) + + def poll(self): + """Poll only from a globally ordered pre-forward boundary.""" + if self.in_flight: + self._poll_in_flight() + return + if self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): + self._poll_evaluation() + + def step(self): + if self.in_flight or self.evaluation_in_flight: + return + self.prefill_steps += 1 + continuous_start = self._continuous_collection_start_step + continuous_end = self._continuous_collection_end_step + if continuous_end is not None: + if continuous_start is not None and self.prefill_steps == continuous_start: + self._set_recording(True) + if self.prefill_steps >= continuous_end: + self._start_evaluation() + return + sampling_interval = self.sampling_interval + phase = self.prefill_steps % sampling_interval + steady_collection_end_step = self._steady_collection_end_step + if steady_collection_end_step is not None: + if self.prefill_steps >= steady_collection_end_step: + self._steady_collection_end_step = None + self._start_evaluation() + return + if sampling_interval == 1: + self._start_evaluation() + return + if phase == sampling_interval - self._steady_sample_window_steps(): + self._arm_steady_collection(self.prefill_steps + self._steady_sample_window_steps()) + + def _set_recording(self, enabled: bool): + for state in self._eplb_states: + state.recording = enabled + + def _reset_recorded_samples(self): + counters = [state.route_counter for state in self._eplb_states] + if counters: + torch._foreach_zero_(counters) + for state in self._eplb_states: + state.recorded_sample_count = 0 + + def _control_count(self, value: int) -> torch.Tensor: + """Return the main-thread-only reusable control collective scalar.""" + return self._control_ready_count.fill_(value) + + def _clear_continuous_collection(self): + self._continuous_collection_start_step = None + self._continuous_collection_end_step = None + + def _steady_sample_window_steps(self) -> int: + return min(EPLB_STEADY_SAMPLE_STEPS, self.sampling_interval) + + def _arm_steady_collection(self, collection_end_step: int): + """Start the fixed sparse window without moving its evaluation boundary.""" + self._reset_recorded_samples() + self._steady_collection_end_step = collection_end_step + self._set_recording(True) + + def _begin_continuous_collection(self): + minimum_end = self.prefill_steps + self.step_interval + collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval + self._reset_recorded_samples() + self._steady_collection_end_step = None + self._continuous_collection_start_step = collection_end - self.step_interval + self._continuous_collection_end_step = collection_end + self._set_recording(self._continuous_collection_start_step == self.prefill_steps) + + def _prepare_next_sampling_window(self): + """Clear the current window and arm the next sparse sampling window.""" + self._clear_continuous_collection() + self._steady_collection_end_step = None + if self.sampling_interval == 1: + self._reset_recorded_samples() + self._set_recording(True) + elif self.sampling_interval <= EPLB_STEADY_SAMPLE_STEPS: + # There is no later pre-boundary manager step at which to arm a + # full clamped window, so arm immediately but keep the same next + # fixed boundary. + self._arm_steady_collection(self.prefill_steps + self.sampling_interval) + else: + self._reset_recorded_samples() + self._set_recording(False) + + @staticmethod + def _recent_ring_samples(counter: torch.Tensor, recorded_sample_count: int) -> torch.Tensor: + """Return the newest ring rows in chronological order.""" + capacity = counter.shape[0] + available = min(recorded_sample_count, capacity) + if available == 0: + return counter[:0] + start = (recorded_sample_count - available) % capacity + indices = (torch.arange(available, dtype=torch.int64, device=counter.device) + start) % capacity + return counter.index_select(0, indices) + + def _collect_local_samples(self) -> torch.Tensor: + counters = [state.route_counter for state in self._eplb_states] + capacities = [counter.shape[0] for counter in counters] + if len(set(capacities)) != 1 or any(counter.ndim != 2 for counter in counters): + raise RuntimeError("EPLB sample capacities differ between layers") + counts = [state.recorded_sample_count for state in self._eplb_states] + if len(set(counts)) != 1: + raise RuntimeError("EPLB recorded sample counts differ between layers") + sample_count = counts[0] + # Validate the metadata before copying the newest rows to the CPU. + metadata = torch.tensor([sample_count, -sample_count, capacities[0], -capacities[0]], dtype=torch.int64) + dist.all_reduce(metadata, op=dist.ReduceOp.MIN, group=self.evaluation_group) + if metadata[0] != -metadata[1] or metadata[2] != -metadata[3]: + raise RuntimeError("EPLB recorded sample count or capacity differs between ranks") + # Stack the fixed-size ring buffers in one GPU launch. Slicing each + # layer before stacking turns a single launch into one index_select per + # MoE layer and is measurably slower in the normal sparse path. + counter_samples = torch.stack(counters, dim=1) + return self._recent_ring_samples(counter_samples, sample_count).cpu() + + def _commit_layer_metadata(self, layer_index: int): + eplb_state = self._eplb_states[layer_index] + logical_to_physical, replica_count = self.target_metadata[layer_index] + eplb_state.logical_to_physical_map.copy_(logical_to_physical, non_blocking=True) + eplb_state.logical_replica_count.copy_(replica_count, non_blocking=True) + + def _finish_rebalance(self): + self.current_placement = self.target_placement + self.target_placement = None + self.target_metadata = None + self.in_flight = False + self._prepare_next_sampling_window() + if self.global_rank == 0: + logger.info(f"eplb completed wall_time={time.time() - self.in_flight_started_at:.2f}s") + + def _poll_in_flight(self): + local_error = None + try: + pending = self.transfer.pending_layers() + except BaseException as exc: + pending = [] + local_error = exc + ready_count = self._control_count(EPLB_CONTROL_ERROR if local_error is not None else len(pending)) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if local_error is not None: + raise RuntimeError("EPLB transfer worker failed on this rank") from local_error + raise RuntimeError("EPLB transfer worker failed on another rank") + if ready_count == 0: + return + if ready_count > len(pending) or ready_count > len(self.in_flight_layers): + raise RuntimeError("EPLB global ready count exceeds the local ordered prefix") + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + # Previous forward is queued on the shared overlap stream; order the + # live-weight commit after it. The subsequent wait orders the next forward. + torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) + for layer_index, buffer_index in pending[:ready_count]: + if layer_index != self.in_flight_layers[0]: + raise RuntimeError( + f"EPLB pending layer {layer_index} does not match expected {self.in_flight_layers[0]}" + ) + self.transfer.commit(layer_index, buffer_index, lambda: self._commit_layer_metadata(layer_index)) + self.in_flight_layers.pop(0) + if not self.in_flight_layers: + self.transfer.finish() + self._finish_rebalance() + + def _plan_and_broadcast(self, global_load: torch.Tensor): + """Plan on rank zero and share the serializable result on the evaluation group.""" + result = None + local_error = None + if self.global_rank == 0: + try: + minimum = self.num_logical_experts * EPLB_MIN_AVG_TOKENS_PER_EXPERT + layer_samples = global_load.sum(dim=(0, 2, 3)) + if torch.any(layer_samples < minimum): + result = { + "kind": "insufficient", + "minimum_layer_samples": int(layer_samples.min().item()), + "minimum": minimum, + } + else: + candidate = plan_redundant_experts( + global_load, + self.world_size, + self.num_redundant_experts_per_rank, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + node_world_size=self.node_world_size, + current_placement=self.current_placement, + stickiness=self.placement_stickiness, + ) + placement, improved, metrics, before_load, after_load = select_improving_placements( + global_load, + self.current_placement, + candidate, + expert_alignment=EPLB_EXPERT_ALIGNMENT, + node_world_size=self.node_world_size, + rebalance_gain_threshold=self.rebalance_gain_threshold, + ) + if bool(torch.any(improved)): + # A planner placement identifies experts by rank, not + # by redundant slot. Canonicalize every selected row + # before broadcasting so transfer, metadata, and the + # next current_placement all describe the same live + # physical expert rows. + placement = placement.clone() + for layer_index in torch.nonzero(improved, as_tuple=False).flatten().tolist(): + placement[layer_index] = align_target_placement( + self.current_placement[layer_index], placement[layer_index] + ) + result = { + "kind": "planned" if bool(torch.any(improved)) else "no_improvement", + "placement": placement, + "improved": improved, + "before": _imbalance_summary(before_load), + "after": _imbalance_summary(after_load), + **metrics, + } + except BaseException as exc: + local_error = exc + result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} + if self.world_size > 1: + result_list = [result] + dist.broadcast_object_list(result_list, src=0, group=self.evaluation_group) + result = result_list[0] + if result["kind"] == "error": + if local_error is not None: + raise RuntimeError("EPLB planner failed on rank zero") from local_error + raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") + return result + + def _evaluate_after_event(self, event: torch.cuda.Event): + """Run the CPU/Gloo planning phase after the frozen CUDA counters are ready.""" + try: + torch.cuda.set_device(self._eplb_states[0].route_counter.device) + event.synchronize() + local_load = self._collect_local_samples() + recorded_sample_count = int(local_load.shape[0]) + sample_window_steps = ( + self.step_interval + if self._continuous_collection_end_step is not None + else self._steady_sample_window_steps() + ) + num_nodes = self.world_size // self.node_world_size + # Preserve source nodes until physical-replica loads are combined; + # DeepEP applies expert alignment after traffic from all sources + # reaches each destination expert. + global_load = torch.zeros((*local_load.shape[:2], num_nodes, local_load.shape[2]), dtype=local_load.dtype) + global_load[:, :, self.global_rank // self.node_world_size] = local_load + dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) + result = self._plan_and_broadcast(global_load) + result["recorded_sample_count"] = recorded_sample_count + result["sample_window_steps"] = sample_window_steps + if result["kind"] == "planned": + metadata = [None] * len(self.weights) + layer_plans = [] + improved_layer_indices = torch.nonzero(result["improved"], as_tuple=False).flatten() + if improved_layer_indices.numel(): + maps_for_improved_layers, counts_for_improved_layers = build_logical_to_physical_maps_for_layers( + result["placement"][improved_layer_indices], + self.num_logical_experts, + source_rank=self.global_rank, + node_world_size=self.node_world_size, + ) + for improved_layer_offset, layer_index in enumerate(improved_layer_indices.tolist()): + placement = result["placement"][layer_index] + metadata[layer_index] = ( + maps_for_improved_layers[improved_layer_offset], + counts_for_improved_layers[improved_layer_offset], + ) + layer_plans.append( + ( + layer_index, + build_transfer_plan( + self.current_placement[layer_index], + placement, + self.num_logical_experts, + self.world_size, + self.node_world_size, + ), + ) + ) + result["metadata"] = metadata + result["layer_plans"] = layer_plans + result["prepared_batches"] = self.transfer.prepare_transfer(layer_plans) + with self._evaluation_lock: + self._evaluation_result = result + except BaseException as exc: + with self._evaluation_lock: + self._evaluation_error = exc + + def _start_evaluation(self): + with self._evaluation_lock: + self._evaluation_result = None + self._evaluation_error = None + self._set_recording(False) + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + self.evaluation_in_flight = True + self._evaluation_thread = threading.Thread(target=self._evaluate_after_event, args=(event,), daemon=True) + self._evaluation_thread.start() + + def _poll_evaluation(self): + if not self.evaluation_in_flight: + return False + with self._evaluation_lock: + error = self._evaluation_error + result = self._evaluation_result + if error is not None or result is not None: + self._evaluation_result = None + self._evaluation_error = None + if error is not None: + self._evaluation_thread.join() + self.evaluation_in_flight = False + self._evaluation_thread = None + raise error + if result is None: + return True + self._evaluation_thread.join() + self.evaluation_in_flight = False + self._evaluation_thread = None + if result["kind"] == "insufficient": + from_continuous_window = self._continuous_collection_end_step is not None + if from_continuous_window: + self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) + self._prepare_next_sampling_window() + else: + self._begin_continuous_collection() + if self.global_rank == 0: + if from_continuous_window: + logger.info( + "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " + "next_sampling_interval=%s recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["minimum_layer_samples"], + result["minimum"], + self.sampling_interval, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + else: + logger.info( + "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " + "scheduled_fresh_window_start=%s scheduled_fresh_window_end=%s " + "recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["minimum_layer_samples"], + result["minimum"], + self._continuous_collection_start_step, + self._continuous_collection_end_step, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + return False + if result["kind"] == "no_improvement": + self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) + if self.global_rank == 0: + logger.info( + "eplb skip rearrangement: no model improvement model_imbalance_ratio=%.4f " + "candidate_model_imbalance_ratio=%.4f candidate_rebalance_gain=%.4f " + "candidate_changed_layer_count=%s actual_changed_layer_count=0 next_sampling_interval=%s " + "recorded_sample_count=%s sample_window_steps=%s", + result["model_imbalance_ratio"], + result["candidate_model_imbalance_ratio"], + result["candidate_rebalance_gain"], + result["candidate_changed_layer_count"], + self.sampling_interval, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + self._prepare_next_sampling_window() + return False + self._start_rebalance(result) + return True + + def _evaluation_ready_on_all_ranks(self) -> bool: + with self._evaluation_lock: + local_error = self._evaluation_error + local_result = self._evaluation_result + local_status = EPLB_CONTROL_ERROR if local_error is not None else int(local_result is not None) + ready_count = self._control_count(local_status) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) + ready_count = int(ready_count.item()) + if ready_count < 0: + if local_error is not None: + raise RuntimeError("EPLB evaluation failed on this rank") from local_error + raise RuntimeError("EPLB evaluation failed on another rank") + return bool(ready_count) + + def _start_rebalance(self, result): + placement = result["placement"] + layer_plans = result["layer_plans"] + self.sampling_interval = self.step_interval + self._clear_continuous_collection() + self._reset_recorded_samples() + self.target_placement = placement + self.target_metadata = result["metadata"] + self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] + self.in_flight = True + self.in_flight_started_at = time.time() + self.transfer.start(layer_plans, result["prepared_batches"]) + if self.global_rank == 0: + actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) + cross_node_transfer_count = sum( + step.src_rank // self.node_world_size != step.dst_rank // self.node_world_size + for _, plan in layer_plans + for step in plan + ) + logger.info( + "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f p95_before=%.4f p95_after=%.4f " + "model_imbalance_ratio=%.4f candidate_model_imbalance_ratio=%.4f " + "candidate_rebalance_gain=%.4f candidate_changed_layer_count=%s " + "actual_changed_layer_count=%s actual_changed_slot_count=%s cross_node_transfer_count=%s " + "recorded_sample_count=%s sample_window_steps=%s", + self.prefill_steps, + result["before"]["max"], + result["after"]["max"], + result["before"]["p95"], + result["after"]["p95"], + result["model_imbalance_ratio"], + result["candidate_model_imbalance_ratio"], + result["candidate_rebalance_gain"], + result["candidate_changed_layer_count"], + len(layer_plans), + actual_changed_slot_count, + cross_node_transfer_count, + result.get("recorded_sample_count"), + result.get("sample_window_steps"), + ) + + +def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: + if rank_load.ndim == 2: + critical = rank_load.max(dim=1).values + mean = rank_load.mean(dim=1) + elif rank_load.ndim == 3: + critical = rank_load.max(dim=2).values.sum(dim=0) + mean = rank_load.mean(dim=2).sum(dim=0) + else: + raise ValueError("rank_load must be [layers, ranks] or [samples, layers, ranks]") + layer_imbalance = critical / mean.clamp_min(1.0) + sorted_imbalance = torch.sort(layer_imbalance).values + p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) + return { + "max": float(layer_imbalance.max().item()), + "p95": float(sorted_imbalance[p95_index].item()), + } + + +def _find_fused_moe_weights(model): + weights_by_id = {} + for layer in model.trans_layers_weight: + for value in getattr(layer, "__dict__", {}).values(): + if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: + weights_by_id[id(value)] = value + return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py new file mode 100644 index 0000000000..9d0c9f120e --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -0,0 +1,775 @@ +"""Asynchronous expert-row migration for EPLB.""" +import ctypes +import os +import re +import socket +import threading +from collections import defaultdict, deque +from dataclasses import dataclass +from typing import Dict, Iterable, List, Optional, Sequence, Tuple + +import torch +import torch.distributed as dist + +from lightllm.common.eplb_utils import EPLB_MAX_STAGING_DEPTH, extract_eplb_expert_tensors + + +@dataclass(frozen=True) +class TransferStep: + dst_rank: int + dst_slot: int + src_rank: int + src_local_row: int + + +def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + """Canonicalize a target row layout without moving retained experts. + + EPLB placement is rank-based: redundant slots on one rank are + interchangeable. Retained experts therefore keep their live physical + slot, while new experts fill freed slots in the planner's target-row + order. The returned placement is the single canonical layout that must + be used both for transfers and for published routing metadata. + """ + assert current.ndim == target.ndim == 2 + assert tuple(current.shape) == tuple(target.shape) + + current_rows = current.tolist() + target_rows = target.tolist() + aligned_target_rows = [] + for current_row, target_row in zip(current_rows, target_rows): + current_slots = {expert: slot for slot, expert in enumerate(current_row)} + target_experts = set(target_row) + aligned_row = list(current_row) + freed_slots = [slot for slot, expert in enumerate(current_row) if expert not in target_experts] + new_experts = [expert for expert in target_row if expert not in current_slots] + assert len(freed_slots) == len(new_experts) + for slot, expert in zip(freed_slots, new_experts): + aligned_row[slot] = expert + aligned_target_rows.append(aligned_row) + return target.new_tensor(aligned_target_rows) + + +def build_transfer_plan( + current: torch.Tensor, + target: torch.Tensor, + num_logical_experts: int, + world_size: int, + node_world_size: int, +) -> List[TransferStep]: + assert tuple(current.shape) == tuple(target.shape) == (world_size, current.shape[1]) + num_experts_per_rank = num_logical_experts // world_size + current_rows = current.tolist() + aligned_target_rows = align_target_placement(current, target).tolist() + # A logical expert has one primary row and at most one redundant row per + # rank, so this source list is already unique. Build it once instead of + # allocating/sorting a set for every destination slot. + candidates_by_expert = [ + [ + ( + expert // num_experts_per_rank, + expert % num_experts_per_rank, + ) + ] + for expert in range(num_logical_experts) + ] + for rank, row in enumerate(current_rows): + for slot, expert in enumerate(row): + candidates_by_expert[expert].append((rank, num_experts_per_rank + slot)) + source_load = [0] * world_size + plan = [] + for dst_rank in range(world_size): + for dst_slot, expert in enumerate(aligned_target_rows[dst_rank]): + if expert == current_rows[dst_rank][dst_slot]: + continue + src_rank, src_row = min( + candidates_by_expert[expert], + key=lambda item: ( + item[0] // node_world_size != dst_rank // node_world_size, + source_load[item[0]], + item[0], + item[1], + ), + ) + source_load[src_rank] += 1 + plan.append(TransferStep(dst_rank, dst_slot, src_rank, src_row)) + return plan + + +class _EPLBTransferBase: + """Shared live/staging buffers and publish/commit lifecycle.""" + + staging_depth = 1 + + def __init__(self, weights, transfer_group, global_rank, world_size): + self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] + self.transfer_group = transfer_group + self.global_rank = global_rank + self.world_size = world_size + self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank + self.device = weights[0].w13.weight.device + self.live = [extract_eplb_expert_tensors(weight) for weight in weights] + self._validate_live_layout(weights) + num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank + self.staging = [ + [ + ( + name, + torch.empty( + (num_redundant_slots_per_rank,) + tuple(tensor.shape[1:]), + dtype=tensor.dtype, + device=tensor.device, + ), + ) + for name, tensor in self.live[0] + ] + for _ in range(self.staging_depth) + ] + self._release = [threading.Event() for _ in range(self.staging_depth)] + for release in self._release: + release.set() + self._error = None + self._consumed_events = [torch.cuda.Event() for _ in range(self.staging_depth)] + self._consumed_recorded = [False] * self.staging_depth + self._changed_dst_slots = [()] * self.staging_depth + self._pending = deque() + self._pending_lock = threading.Lock() + self._thread = None + self._needs_staging_reuse_barrier = False + + def _validate_live_layout(self, weights) -> None: + reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] + num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank + for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): + layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] + assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" + assert ( + state.num_redundant_experts_per_rank == num_redundant_slots_per_rank + ), "EPLB redundant slot count must match" + + def _copy_batch(self, batch, prepared_batch) -> None: + raise NotImplementedError + + def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [ + [ + (layer_index, plan, buffer_index, self.staging[buffer_index]) + for buffer_index, (layer_index, plan) in enumerate( + layer_plans[batch_start : batch_start + self.staging_depth] + ) + ] + for batch_start in range(0, len(layer_plans), self.staging_depth) + ] + + def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [(batch, None) for batch in self._make_batches(layer_plans)] + + def _start_transfer_generation(self) -> None: + """Prepare backend state after the in-flight worker check succeeds.""" + + def _finish_transfer_generation(self) -> None: + """Release backend state only after the migration worker has joined.""" + + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches=None) -> None: + if self._thread is not None and self._thread.is_alive(): + raise RuntimeError("EPLB transfer is already in flight") + if prepared_batches is None: + prepared_batches = self.prepare_transfer(layer_plans) + expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth + if len(prepared_batches) != expected_batch_count: + raise ValueError("EPLB prepared batch count does not match layer-plan batches") + # 预构造批次与描述符一起传入,避免推理线程重建。 + self._start_transfer_generation() + self._error = None + with self._pending_lock: + self._pending.clear() + + def worker() -> None: + try: + torch.cuda.set_device(self.device) + if not layer_plans: + self._finish_transfer_generation() + for batch_index, (batch, prepared_batch) in enumerate(prepared_batches): + batch_start = batch_index * self.staging_depth + for layer_index, plan, buffer_index, _ in batch: + release = self._release[buffer_index] + # A buffer cannot be reused until its prior committed rows are no longer read by CUDA. + release.wait() + release.clear() + if self._consumed_recorded[buffer_index]: + self._consumed_events[buffer_index].synchronize() + self._changed_dst_slots[buffer_index] = tuple( + step.dst_slot for step in plan if step.dst_rank == self.global_rank + ) + if batch_start > 0 and self._needs_staging_reuse_barrier: + # All destinations must finish consuming the prior IPC staging generation + # before a source can reuse the peer buffer for this batch. + dist.barrier(group=self.transfer_group) + self._copy_batch(batch, prepared_batch) + if batch_start + self.staging_depth >= len(layer_plans): + self._finish_transfer_generation() + with self._pending_lock: + self._pending.extend((layer_index, buffer_index) for layer_index, _, buffer_index, _ in batch) + except BaseException as exc: + self._error = exc + + self._thread = threading.Thread(target=worker, name=f"eplb-{self.backend}", daemon=True) + self._thread.start() + + def pending_layers(self): + if self._error is not None: + raise RuntimeError("EPLB migration worker failed") from self._error + with self._pending_lock: + return list(self._pending) + + def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: + with self._pending_lock: + if not self._pending or self._pending[0] != (layer_index, buffer_index): + raise RuntimeError("EPLB commit does not match the pending FIFO") + self._pending.popleft() + changed_dst_slots = self._changed_dst_slots[buffer_index] + for (_, live), (_, staging) in zip(self.live[layer_index], self.staging[buffer_index]): + _commit_staging_rows( + live, + staging, + self.num_experts_per_rank, + changed_dst_slots, + ) + if post_copy is not None: + post_copy() + self._consumed_events[buffer_index].record(torch.cuda.current_stream()) + self._consumed_recorded[buffer_index] = True + self._release[buffer_index].set() + + def finish(self) -> None: + """Wait for the released migration worker to exit before another rebalance.""" + thread = self._thread + if thread is None: + return + thread.join() + self._thread = None + if self._error is not None: + raise RuntimeError("EPLB migration worker failed") from self._error + + +class NixlEPLBTransfer(_EPLBTransferBase): + """GPU-direct UCX/NIXL EPLB transfer. Initialization errors are fatal.""" + + backend = "nixl" + _DEFAULT_UCX_TLS = "self,sm,cuda_ipc,cuda_copy,rc_x" + + @dataclass + class _PreparedBatch: + remote_entries: Dict[int, list] + push_batch: "_PreparedCudaMemcpyBatch | None" + + def __init__(self, weights, transfer_group, global_rank, world_size): + # Reuse at most eight layer buffers to bound EPLB staging memory. + self.staging_depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) + super().__init__(weights, transfer_group, global_rank, world_size) + self._nixl_agent = None + self._registered_descs = None + self._remote_agents: Dict[int, str] = {} + self._remote_layouts = {} + self._xfer_cache = {} + self._used_xfer_cache_keys = set() + self._ipc_staging = {} + self._same_node_ranks = set() + self._cross_node_ranks = set() + self._push_stream = torch.cuda.Stream(device=self.device) + self._batch_memcpy = _CudaBatchMemcpy() + try: + self._init_ipc_metadata() + self._init_push_layouts() + if self._cross_node_ranks: + os.environ.setdefault("UCX_TLS", self._DEFAULT_UCX_TLS) + try: + import nixl + except Exception as exc: + raise RuntimeError("NIXL EPLB backend requires the nixl package for cross-node transfer") from exc + agent_name = f"lightllm-eplb-{socket.gethostname()}-{os.getpid()}-rank-{global_rank}" + config = nixl.nixl_agent_config(enable_prog_thread=True, enable_listen_thread=False, backends=["UCX"]) + self._nixl_agent = nixl.nixl_agent(agent_name, config) + reg_tensors = [tensor for layer in self.live for _, tensor in layer] + [ + tensor for staging in self.staging for _, tensor in staging + ] + self._registered_descs = self._nixl_agent.get_reg_descs(reg_tensors) + self._nixl_agent.register_memory(self._registered_descs, backends=["UCX"]) + self._init_remote_metadata() + except Exception as exc: + self.shutdown() + if isinstance(exc, RuntimeError): + raise + raise RuntimeError("NIXL EPLB initialization failed") from exc + + def _local_layout(self): + return [ + [(name, tensor.data_ptr(), tensor.get_device(), tensor[0].nbytes) for name, tensor in layer] + for layer in self.live + ] + + def _init_ipc_metadata(self) -> None: + hostnames = [None] * self.world_size + dist.all_gather_object(hostnames, socket.gethostname(), group=self.transfer_group) + local_hostname = hostnames[self.global_rank] + self._needs_staging_reuse_barrier = len(set(hostnames)) < len(hostnames) + self._same_node_ranks = {rank for rank, hostname in enumerate(hostnames) if hostname == local_hostname} + self._cross_node_ranks = set(range(self.world_size)) - self._same_node_ranks + from lightllm.server.router.model_infer.mode_backend.pd.p2p_fix import ( + p2p_fix_rebuild_cuda_tensor, + reduce_tensor, + ) + + exports = {} + for target_rank in self._same_node_ranks - {self.global_rank}: + exports[target_rank] = { + "staging": [ + [(name, tuple(tensor.shape), tensor.dtype, reduce_tensor(tensor)[1]) for name, tensor in staging] + for staging in self.staging + ], + } + all_exports = [None] * self.world_size + dist.all_gather_object(all_exports, exports, group=self.transfer_group) + + torch.cuda.set_device(self.device) + for dst_rank in self._same_node_ranks - {self.global_rank}: + metadata = all_exports[dst_rank].get(self.global_rank) + if metadata is None or len(metadata["staging"]) != self.staging_depth: + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} has incompatible staging metadata") + rebuilt_staging = [] + for remote_staging, local_staging in zip(metadata["staging"], self.staging): + if len(remote_staging) != len(local_staging): + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging tensor count mismatch") + rebuilt = [] + for (name, shape, dtype, args), (local_name, local_tensor) in zip(remote_staging, local_staging): + if name != local_name or shape != tuple(local_tensor.shape) or dtype != local_tensor.dtype: + raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging layout mismatch for {name}") + tensor = p2p_fix_rebuild_cuda_tensor(*args) + if tuple(tensor.shape) != shape or tensor.dtype != dtype or tensor.device != local_tensor.device: + raise RuntimeError( + f"NIXL IPC destination rank {dst_rank} staging rebuild validation failed for {name}" + ) + rebuilt.append((name, tensor)) + rebuilt_staging.append(rebuilt) + self._ipc_staging[dst_rank] = rebuilt_staging + + def _init_remote_metadata(self) -> None: + metadata = self._nixl_agent.get_agent_metadata() + all_metadata = [None] * self.world_size + all_layouts = [None] * self.world_size + dist.all_gather_object(all_metadata, metadata, group=self.transfer_group) + dist.all_gather_object(all_layouts, self._local_layout(), group=self.transfer_group) + for rank in self._cross_node_ranks: + layout = all_layouts[rank] + if len(layout) != len(self.live): + raise RuntimeError(f"NIXL remote rank {rank} has incompatible layer layout") + self._remote_agents[rank] = self._nixl_agent.add_remote_agent(all_metadata[rank]) + self._remote_layouts[rank] = layout + + def _wait_xfers(self, xfers) -> None: + pending = [] + for item in xfers: + state = self._nixl_agent.transfer(item[2]) + if state == "ERR": + raise RuntimeError("NIXL READ post failed") + if state == "PROC": + pending.append(item) + while pending: + remaining = [] + for item in pending: + state = self._nixl_agent.check_xfer_state(item[2]) + if state == "ERR": + raise RuntimeError("NIXL READ transfer failed") + if state != "DONE": + remaining.append(item) + pending = remaining + + def _release_xfers(self, xfers) -> None: + unreleased = [] + errors = [] + for local_dlist, remote_dlist, xfer in xfers: + remaining = [local_dlist, remote_dlist, xfer] + for remaining_index, handle, release in ( + (2, xfer, self._nixl_agent.release_xfer_handle), + (1, remote_dlist, self._nixl_agent.release_dlist_handle), + (0, local_dlist, self._nixl_agent.release_dlist_handle), + ): + if handle is not None: + try: + release(handle) + except Exception as exc: + errors.append(exc) + else: + remaining[remaining_index] = None + if any(handle is not None for handle in remaining): + unreleased.append(tuple(remaining)) + if errors: + error = RuntimeError("NIXL transfer handle release failed") + error.unreleased_xfers = unreleased + raise error from errors[0] + + @staticmethod + def _contiguous_runs(steps): + ordered = sorted(steps, key=lambda step: (step.src_local_row, step.dst_slot)) + runs = [] + for step in ordered: + if ( + runs + and step.src_local_row == runs[-1][-1].src_local_row + 1 + and step.dst_slot == runs[-1][-1].dst_slot + 1 + ): + runs[-1].append(step) + else: + runs.append([step]) + return runs + + @staticmethod + def _remote_read_cache_key(src_rank: int, entries): + return ( + src_rank, + tuple( + ( + layer_index, + tuple((step.src_local_row, step.dst_slot) for step in run), + tuple(tensor.data_ptr() for _, tensor in staging), + ) + for layer_index, run, staging in entries + ), + ) + + def _init_push_layouts(self) -> None: + self._live_row_layout = [ + [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in layer] for layer in self.live + ] + reference = [(name, row_nbytes) for name, _, row_nbytes in self._live_row_layout[0]] + self._push_staging_row_layout = {} + for dst_rank in self._same_node_ranks: + layouts = [] + for buffer_index in range(self.staging_depth): + staging = ( + self.staging[buffer_index] + if dst_rank == self.global_rank + else self._ipc_staging[dst_rank][buffer_index] + ) + layout = [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in staging] + if [(name, row_nbytes) for name, _, row_nbytes in layout] != reference: + raise RuntimeError("NIXL source-push staging row layout mismatch") + layouts.append(layout) + self._push_staging_row_layout[dst_rank] = layouts + + def _prepare_batch(self, batch): + remote_entries = defaultdict(list) + push_descriptors = [] + for layer_index, plan, buffer_index, staging in batch: + steps_by_source = defaultdict(list) + by_destination = defaultdict(list) + for step in plan: + if step.dst_rank == self.global_rank and step.src_rank not in self._same_node_ranks: + steps_by_source[step.src_rank].append(step) + if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: + by_destination[step.dst_rank].append(step) + for src_rank, steps in steps_by_source.items(): + remote_entries[src_rank].extend((layer_index, run, staging) for run in self._contiguous_runs(steps)) + source_layout = self._live_row_layout[layer_index] + for dst_rank, steps in by_destination.items(): + destination_layout = self._push_staging_row_layout[dst_rank][buffer_index] + for run in self._contiguous_runs(steps): + first = run[0] + run_len = len(run) + for (_, source_ptr, row_nbytes), (_, destination_ptr, _) in zip(source_layout, destination_layout): + push_descriptors.append( + ( + source_ptr + first.src_local_row * row_nbytes, + destination_ptr + first.dst_slot * row_nbytes, + run_len * row_nbytes, + ) + ) + push_batch = self._batch_memcpy.prepare(push_descriptors) if push_descriptors else None + return self._PreparedBatch(dict(remote_entries), push_batch) + + def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): + return [(batch, self._prepare_batch(batch)) for batch in self._make_batches(layer_plans)] + + def _get_remote_read(self, src_rank: int, entries): + cache_key = self._remote_read_cache_key(src_rank, entries) + cached = self._xfer_cache.get(cache_key) + if cached is not None: + self._used_xfer_cache_keys.add(cache_key) + return cached + local_descs = [] + remote_descs = [] + local_dlist = remote_dlist = xfer = None + try: + for layer_index, run, staging in entries: + remote_layer = self._remote_layouts[src_rank][layer_index] + if len(remote_layer) != len(staging): + raise RuntimeError(f"NIXL remote rank {src_rank} has incompatible layer layout") + first = run[0] + run_len = len(run) + for tensor_index, (_, staging_tensor) in enumerate(staging): + name, remote_ptr, remote_device, remote_nbytes = remote_layer[tensor_index] + if ( + name != self.live[layer_index][tensor_index][0] + or remote_nbytes != staging_tensor[first.dst_slot].nbytes + ): + raise RuntimeError(f"NIXL remote rank {src_rank} descriptor range mismatch") + local_descs.append( + ( + staging_tensor[first.dst_slot].data_ptr(), + run_len * remote_nbytes, + staging_tensor.get_device(), + ) + ) + remote_descs.append( + (remote_ptr + first.src_local_row * remote_nbytes, run_len * remote_nbytes, remote_device) + ) + local_dlist = self._nixl_agent.prep_xfer_dlist( + "NIXL_INIT_AGENT", self._nixl_agent.get_xfer_descs(local_descs, "VRAM"), backends=["UCX"] + ) + remote_dlist = self._nixl_agent.prep_xfer_dlist( + self._remote_agents[src_rank], self._nixl_agent.get_xfer_descs(remote_descs, "VRAM"), backends=["UCX"] + ) + xfer = self._nixl_agent.make_prepped_xfer( + "READ", + local_dlist, + list(range(len(local_descs))), + remote_dlist, + list(range(len(remote_descs))), + backends=["UCX"], + ) + selected_backend = self._nixl_agent.query_xfer_backend(xfer) + if selected_backend != "UCX": + raise RuntimeError("NIXL EPLB READ did not select UCX") + self._xfer_cache[cache_key] = (local_dlist, remote_dlist, xfer) + self._used_xfer_cache_keys.add(cache_key) + return self._xfer_cache[cache_key] + except Exception: + self._release_xfers([(local_dlist, remote_dlist, xfer)]) + raise + + def _copy_batch(self, batch, prepared_batch) -> None: + if prepared_batch.push_batch is not None: + self._batch_memcpy.enqueue(prepared_batch.push_batch, self._push_stream.cuda_stream) + xfers = [ + self._get_remote_read(src_rank, entries) for src_rank, entries in prepared_batch.remote_entries.items() + ] + self._wait_xfers(xfers) + self._push_stream.synchronize() + # Before a rank publishes this batch it has completed its outgoing source-pushes and + # incoming UCX READs. The manager's global MIN-ready gate therefore means all transfers + # are complete before any rank commits, without a destination-side GPU wait. + + def _start_transfer_generation(self) -> None: + self._used_xfer_cache_keys.clear() + + def _finish_transfer_generation(self) -> None: + errors = [] + for cache_key in set(self._xfer_cache) - self._used_xfer_cache_keys: + xfer = self._xfer_cache[cache_key] + try: + self._release_xfers([xfer]) + except Exception as exc: + unreleased = getattr(exc, "unreleased_xfers", None) + if unreleased: + self._xfer_cache[cache_key] = unreleased[0] + errors.append(exc) + else: + del self._xfer_cache[cache_key] + if errors: + raise RuntimeError("NIXL EPLB cache eviction failed") from errors[0] + + def shutdown(self) -> None: + agent = self._nixl_agent + errors = [] + getattr(self, "_used_xfer_cache_keys", set()).clear() + if agent is not None: + for cache_key, xfer in list(self._xfer_cache.items()): + try: + self._release_xfers([xfer]) + except Exception as exc: + unreleased = getattr(exc, "unreleased_xfers", None) + if unreleased: + self._xfer_cache[cache_key] = unreleased[0] + errors.append(exc) + else: + del self._xfer_cache[cache_key] + if errors: + raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] + for remote_name in list(self._remote_agents.values()): + if agent is not None: + try: + agent.remove_remote_agent(remote_name) + except Exception as exc: + errors.append(exc) + self._remote_agents.clear() + self._remote_layouts.clear() + if agent is not None and self._registered_descs is not None: + try: + agent.deregister_memory(self._registered_descs, backends=["UCX"]) + except Exception as exc: + errors.append(exc) + self._registered_descs = None + self._nixl_agent = None + getattr(self, "_ipc_staging", {}).clear() + if errors: + raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] + + def __del__(self): + try: + self.shutdown() + except Exception: + pass + + +def _commit_staging_rows( + live: torch.Tensor, + staging: torch.Tensor, + num_experts_per_rank: int, + changed_dst_slots: Sequence[int], +) -> None: + slots = sorted(set(changed_dst_slots)) + if not slots: + return + run_start = previous = slots[0] + for dst_slot in (*slots[1:], None): + if dst_slot is not None and dst_slot == previous + 1: + previous = dst_slot + continue + run_length = previous - run_start + 1 + live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( + staging.narrow(0, run_start, run_length), non_blocking=True + ) + if dst_slot is not None: + run_start = previous = dst_slot + + +class _CudaMemLocation(ctypes.Structure): + _fields_ = [("type", ctypes.c_int), ("id", ctypes.c_int)] + + +class _CudaMemcpyAttributes(ctypes.Structure): + _fields_ = [ + ("srcAccessOrder", ctypes.c_int), + ("srcLocHint", _CudaMemLocation), + ("dstLocHint", _CudaMemLocation), + ("flags", ctypes.c_uint), + ] + + +@dataclass +class _PreparedCudaMemcpyBatch: + """Host-side arrays retained for one cudaMemcpyBatchAsync submission.""" + + dsts: object + srcs: object + sizes: object + attrs: _CudaMemcpyAttributes + attrs_idxs: object + count: int + + +class _CudaBatchMemcpy: + """CUDA 13.x ``cudaMemcpyBatchAsync`` binding for EPLB source-push.""" + + _SRC_ACCESS_ORDER_STREAM = 1 + _PREFER_OVERLAP_WITH_COMPUTE = 1 + _CUDA_13_0 = 13000 + _CUDA_14_0 = 14000 + + def __init__(self, library=None): + if library is None: + path = self._find_loaded_cudart() + if path is None: + raise RuntimeError( + "NIXL same-node source-push requires CUDA Runtime 13.x cudaMemcpyBatchAsync; " + "libcudart.so.13 is not loaded" + ) + try: + library = ctypes.CDLL(path) + except OSError as exc: + raise RuntimeError(f"cannot load libcudart: {exc}") from exc + + try: + runtime_get_version = library.cudaRuntimeGetVersion + self._batch_async = library.cudaMemcpyBatchAsync + self._get_error_string = library.cudaGetErrorString + except AttributeError as exc: + raise RuntimeError("cudaMemcpyBatchAsync is unavailable") from exc + + runtime_get_version.restype = ctypes.c_int + runtime_get_version.argtypes = [ctypes.POINTER(ctypes.c_int)] + self._get_error_string.restype = ctypes.c_char_p + self._get_error_string.argtypes = [ctypes.c_int] + runtime_version = ctypes.c_int() + result = runtime_get_version(ctypes.byref(runtime_version)) + if result != 0: + raise RuntimeError(f"cudaRuntimeGetVersion failed with CUDA error {result}") + if not self._CUDA_13_0 <= runtime_version.value < self._CUDA_14_0: + raise RuntimeError( + f"cudaMemcpyBatchAsync requires CUDA Runtime 13.x (13.0 ABI), found {runtime_version.value}" + ) + + pointer_array = ctypes.POINTER(ctypes.c_void_p) + self._batch_async.restype = ctypes.c_int + self._batch_async.argtypes = [ + pointer_array, + pointer_array, + ctypes.POINTER(ctypes.c_size_t), + ctypes.c_size_t, + ctypes.POINTER(_CudaMemcpyAttributes), + ctypes.POINTER(ctypes.c_size_t), + ctypes.c_size_t, + ctypes.c_void_p, + ] + + @staticmethod + def prepare(copies: Iterable[Tuple[int, int, int]]) -> _PreparedCudaMemcpyBatch: + copies = tuple(copies) + if not copies: + raise ValueError("cudaMemcpyBatchAsync requires at least one copy") + for src, dst, size in copies: + if not src or not dst or size <= 0: + raise ValueError("cudaMemcpyBatchAsync requires non-null pointers and positive sizes") + count = len(copies) + dsts = (ctypes.c_void_p * count)(*(dst for _, dst, _ in copies)) + srcs = (ctypes.c_void_p * count)(*(src for src, _, _ in copies)) + sizes = (ctypes.c_size_t * count)(*(size for _, _, size in copies)) + attrs = _CudaMemcpyAttributes() + attrs.srcAccessOrder = _CudaBatchMemcpy._SRC_ACCESS_ORDER_STREAM + attrs.flags = _CudaBatchMemcpy._PREFER_OVERLAP_WITH_COMPUTE + attrs_idxs = (ctypes.c_size_t * 1)(0) + return _PreparedCudaMemcpyBatch(dsts, srcs, sizes, attrs, attrs_idxs, count) + + def enqueue(self, prepared: _PreparedCudaMemcpyBatch, stream: int) -> None: + result = self._batch_async( + prepared.dsts, + prepared.srcs, + prepared.sizes, + prepared.count, + ctypes.byref(prepared.attrs), + prepared.attrs_idxs, + 1, + ctypes.c_void_p(stream), + ) + if result != 0: + message = self._get_error_string(result) + error = message.decode("utf-8") if message else f"CUDA error {result}" + raise RuntimeError(f"cudaMemcpyBatchAsync failed: {error}") + + @staticmethod + def _find_loaded_cudart() -> Optional[str]: + """Return a mapped CUDA 13 runtime without loading CUDA as a side effect.""" + try: + with open("/proc/self/maps") as maps: + for line in maps: + if "libcudart" not in line: + continue + path_start = line.find("/") + if path_start < 0: + continue + path = line[path_start:].strip().removesuffix(" (deleted)") + if re.search(r"libcudart[^/]*\.so\.13(?:\D|$)", path): + return path + except OSError: + pass + return None diff --git a/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py b/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py deleted file mode 100644 index 596eca4f24..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py +++ /dev/null @@ -1,158 +0,0 @@ -# 对于 deepseekv3 模型在 ep 运行模式下,自动分析统计各个专家的出现频率,然后 -# 自动更新当前的冗余专家为新的冗余专家。 -import torch -import time -import enum -import lightllm.utils.petrel_helper as utils -import threading -import json -from typing import List -from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_redundancy import ( - FusedMoeWeightEPAutoRedundancy, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.utils.envs_utils import get_env_start_args, get_redundancy_expert_update_interval -from lightllm.utils.envs_utils import get_redundancy_expert_update_max_load_count -from lightllm.utils.envs_utils import get_redundancy_expert_num -from lightllm.utils.dist_utils import get_global_rank -from lightllm.common.basemodel.layer_weights.hf_load_utils import load_func -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - - -class RedundancyExpertManager: - def __init__(self, model: TpPartBaseModel): - self.args = get_env_start_args() - self.model = model - self.ep_fused_moeweights: List[FusedMoeWeightEPAutoRedundancy] = [] - for layer in self.model.trans_layers_weight: - ep_weights = self._find_members_of_class(layer, FusedMoeWeight) - assert len(ep_weights) <= 1 - self.ep_fused_moeweights.extend([FusedMoeWeightEPAutoRedundancy(e) for e in ep_weights]) - - # save load params - self.use_safetensors = True - files = utils.PetrelHelper.list(self.args.model_dir, extension="all") - candidate_files = list(filter(lambda x: x.endswith(".safetensors"), files)) - if len(candidate_files) == 0: - self.use_safetensors = False - candidate_files = list(filter(lambda x: x.endswith(".bin"), files)) - assert len(candidate_files) != 0, "can only support pytorch tensor and safetensors format for weights." - self.candidate_files = candidate_files - - # state 1. check_to_update 2. prepare_update 3. start_load_hf_weights 4. wait_load_ready, 5. commit - self.state: _STATE = _STATE.CHECK_TO_UPDATE - self.update_time = time.time() - self.update_interval = get_redundancy_expert_update_interval() - self.load_thread: threading.Thread = None - self.global_rank = get_global_rank() - # 冗余专家的最大加载次数 - self.load_count = 0 - self.max_load_count = get_redundancy_expert_update_max_load_count() - - # 清理counter - self._clear_all_counter() - - self.rank0_redundancy_expert_config = { - "redundancy_expert_num": get_redundancy_expert_num(), - "default": list(range(get_redundancy_expert_num())), - } - - def step(self): - if self.load_count >= self.max_load_count: - return - - if self.state == _STATE.CHECK_TO_UPDATE: - cur_time = time.time() - if cur_time - self.update_time > self.update_interval: - self.update_time = cur_time - self.state = _STATE.PREPARE_UPDATE - logger.info(f"global_rank {self.global_rank} state to prepare update") - elif self.state == _STATE.PREPARE_UPDATE: - self._prepare_load_new_redundancy_expert() - self.state = _STATE.START_LOAD_HF_WEIGHTS - logger.info(f"global_rank {self.global_rank} state to start load hf weights") - - elif self.state == _STATE.START_LOAD_HF_WEIGHTS: - self.load_thread = threading.Thread(target=self._load_hf_weights, daemon=True) - self.load_thread.start() - self.state = _STATE.WAIT_LOAD_READY - logger.info(f"global_rank {self.global_rank} state to wait load ready") - - elif self.state == _STATE.WAIT_LOAD_READY: - if not self.load_thread.is_alive(): - self.load_thread = None - self.state = _STATE.COMMIT - logger.info(f"global_rank {self.global_rank} state to commit") - - elif self.state == _STATE.COMMIT: - self._commit() - self.state = _STATE.CHECK_TO_UPDATE - self.load_count += 1 - logger.info(f"global_rank {self.global_rank} state to check to update") - return - - def _prepare_load_new_redundancy_expert(self): - for w in self.ep_fused_moeweights: - topk_redundancy_expert_ids = w.prepare_redundancy_experts() - if self.global_rank == 0: - self.rank0_redundancy_expert_config[str(w._ep_w.layer_num)] = topk_redundancy_expert_ids - - if self.global_rank == 0: - try: - with open("./redundancy_expert_config.json", "w") as f: - json.dump(self.rank0_redundancy_expert_config, f, indent=4) - logger.info( - f"rank {self.global_rank} save redundancy_expert_config.json to ./redundancy_expert_config.json" - ) - except BaseException as e: - logger.exception(str(e)) - logger.error(f"global rank {self.global_rank} save redundancy_expert_config.json failed") - - return - - def _load_hf_weights(self): - start = time.time() - try: - for file in self.candidate_files: - load_func( - file, - use_safetensors=self.use_safetensors, - pre_post_layer=None, - transformer_layer_list=self.ep_fused_moeweights, - weight_dir=self.args.model_dir, - ) - except BaseException as e: - logger.exception(str(e)) - raise e - cost_time = time.time() - start - logger.info(f"global rank {self.global_rank} load redundancy_expert cost time: {cost_time} s") - return - - def _commit(self): - for w in self.ep_fused_moeweights: - w.commit() - return - - def _find_members_of_class(self, obj, cls): - members = [] - for attr in dir(obj): - value = getattr(obj, attr) - if isinstance(value, cls): - members.append(value) - return members - - def _clear_all_counter(self): - for w in self.ep_fused_moeweights: - w.clear_counter() - return - - -class _STATE(enum.Enum): - CHECK_TO_UPDATE = 0 - PREPARE_UPDATE = 1 - START_LOAD_HF_WEIGHTS = 2 - WAIT_LOAD_READY = 3 - COMMIT = 4 diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 04c78e74ab..5851e6c67b 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -25,7 +25,6 @@ PDDecodeNode, PDDPForDecodeNode, ) -from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( EPBalanceMonitor, @@ -102,13 +101,6 @@ def exposed_init_model(self, kvargs): self.backend.init_model(kvargs) self.rl_backend_ops = RlBackendOps(self.backend) if self.args.enable_rl else None - # only deepseekv3 can support auto_update_redundancy_expert - if self.args.auto_update_redundancy_expert: - self.redundancy_expert_manager = RedundancyExpertManager(self.backend.model) - logger.info("init redundancy_expert_manager") - else: - self.redundancy_expert_manager = None - if should_enable_ep_balance_monitor(self.args): monitor = EPBalanceMonitor(self.backend.model) if monitor.enabled: diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index ff78c95260..b9b10923b6 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -140,67 +140,37 @@ def get_lightllm_websocket_max_message_size(): return int(os.getenv("LIGHTLLM_WEBSOCKET_MAX_SIZE", 128 * 1024 * 1024)) -# get_redundancy_expert_ids and get_redundancy_expert_num are primarily -# used to obtain the IDs and number of redundant experts during inference. -# They depend on a configuration file specified by ep_redundancy_expert_config_path, -# which is a JSON formatted text file. -# The content format is as follows: -# { -# "redundancy_expert_num": 1, # Number of redundant experts per rank -# "0": [0], # Key: layer_index (string), -# # Value: list of original expert IDs that are redundant for this layer -# "1": [0], -# "default": [0] # Default list of redundant expert IDs if layer-specific entry is not found -# } - - -@lru_cache(maxsize=None) -def get_redundancy_expert_ids(layer_index: int): - """ - Get the redundancy expert ids from the environment variable. - :return: List of redundancy expert ids. - """ - args = get_env_start_args() - if args.ep_redundancy_expert_config_path is None: - return [] - - with open(args.ep_redundancy_expert_config_path, "r") as f: - config = json.load(f) - if str(layer_index) in config: - return config[str(layer_index)] - else: - return config.get("default", []) - - @lru_cache(maxsize=None) -def get_redundancy_expert_num(): - """ - Get the number of redundancy experts from the environment variable. - :return: Number of redundancy experts. - """ - args = get_env_start_args() - if args.ep_redundancy_expert_config_path is None: - return 0 - - with open(args.ep_redundancy_expert_config_path, "r") as f: - config = json.load(f) - if "redundancy_expert_num" in config: - return config["redundancy_expert_num"] - else: - return 0 +def get_prefill_eplb_step_interval(): + """Return the number of prefill forwards between EPLB attempts.""" + interval = int(os.getenv("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL", 20)) + if interval <= 0: + raise ValueError("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL must be greater than 0") + return interval @lru_cache(maxsize=None) -def get_redundancy_expert_update_interval(): - return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_INTERVAL", 30 * 60)) +def get_eplb_rebalance_gain_threshold() -> float: + """Return the EPLB gain threshold: estimated critical-load reduction ratio; 0.05 means 5%.""" + env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" + raw_value = os.getenv(env_name, "0.05") + value = float(raw_value) + if not 0.0 <= value <= 1.0: + raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") + return value @lru_cache(maxsize=None) -def get_redundancy_expert_update_max_load_count(): - return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_MAX_LOAD_COUNT", 1)) +def get_eplb_placement_stickiness() -> float: + """Return the EPLB placement stickiness: a keep-bonus, as a fraction of the mean per-layer expert load.""" + env_name = "LIGHTLLM_EPLB_PLACEMENT_STICKINESS" + raw_value = os.getenv(env_name, "0.1") + value = float(raw_value) + if not 0.0 <= value <= 1.0: + raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") + return value -@lru_cache(maxsize=None) def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index 5ec2a9145e..73f172d908 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -74,6 +74,14 @@ def profile_mtp_weight_memory(model): weight_memory_before = torch.cuda.memory_allocated() yield target_weight_bytes = torch.cuda.memory_allocated() - weight_memory_before + # Draft construction can opt out of portions of the target's weight layout + # which it never instantiates (for example EPLB redundant rows). + excluded_weight_bytes = int(getattr(model, "get_mtp_profile_weight_exclusion", lambda: 0)()) + if not 0 <= excluded_weight_bytes <= target_weight_bytes: + raise ValueError( + f"invalid MTP profile exclusion {excluded_weight_bytes}; measured target weights={target_weight_bytes}" + ) + target_weight_bytes -= excluded_weight_bytes model.mem_fraction = get_mtp_adjusted_mem_fraction( mem_fraction=model.mem_fraction, target_weight_bytes=target_weight_bytes, diff --git a/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json b/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json deleted file mode 100644 index 241ab25ea3..0000000000 --- a/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json +++ /dev/null @@ -1,180 +0,0 @@ -{ - "redundancy_expert_num": 1, - "default": [ - 0 - ], - "3": [ - 226 - ], - "4": [ - 123 - ], - "5": [ - 187 - ], - "6": [ - 138 - ], - "7": [ - 132 - ], - "8": [ - 240 - ], - "9": [ - 4 - ], - "10": [ - 88 - ], - "11": [ - 60 - ], - "12": [ - 161 - ], - "13": [ - 178 - ], - "14": [ - 80 - ], - "15": [ - 144 - ], - "16": [ - 195 - ], - "17": [ - 251 - ], - "18": [ - 226 - ], - "19": [ - 87 - ], - "20": [ - 149 - ], - "21": [ - 45 - ], - "22": [ - 214 - ], - "23": [ - 41 - ], - "24": [ - 46 - ], - "25": [ - 156 - ], - "26": [ - 112 - ], - "27": [ - 185 - ], - "28": [ - 58 - ], - "29": [ - 156 - ], - "30": [ - 147 - ], - "31": [ - 199 - ], - "32": [ - 16 - ], - "33": [ - 188 - ], - "34": [ - 227 - ], - "35": [ - 136 - ], - "36": [ - 84 - ], - "37": [ - 15 - ], - "38": [ - 204 - ], - "39": [ - 96 - ], - "40": [ - 226 - ], - "41": [ - 25 - ], - "42": [ - 69 - ], - "43": [ - 122 - ], - "44": [ - 152 - ], - "45": [ - 113 - ], - "46": [ - 98 - ], - "47": [ - 68 - ], - "48": [ - 13 - ], - "49": [ - 102 - ], - "50": [ - 214 - ], - "51": [ - 201 - ], - "52": [ - 182 - ], - "53": [ - 235 - ], - "54": [ - 162 - ], - "55": [ - 125 - ], - "56": [ - 62 - ], - "57": [ - 121 - ], - "58": [ - 105 - ], - "59": [ - 236 - ], - "60": [ - 117 - ] -} \ No newline at end of file diff --git a/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py b/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py deleted file mode 100644 index 16131ef935..0000000000 --- a/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py +++ /dev/null @@ -1,151 +0,0 @@ -import torch -import pytest -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair -from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import expert_id_counter -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) - - -def test_redundancy_topk_ids_repair(): - ep_expert_num = 4 - global_rank = 0 - redundancy_expert_num = 1 - topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - - redundancy_expert_ids = torch.tensor( - [ - 0, - ], - dtype=torch.int64, - device="cuda", - ) - - expert_id_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=redundancy_expert_ids, - ep_expert_num=ep_expert_num, - global_rank=global_rank, - expert_counter=expert_id_counter, - enable_counter=True, - ) - - ans_topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids - new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids - ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( - (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 - ) - - assert torch.equal(topk_ids, ans_topk_ids) - assert torch.equal( - expert_id_counter, torch.tensor([1, 2, 1, 2, 0, 1, 0, 2, 0, 1, 1, 1], dtype=torch.int64, device="cuda") - ) - - ep_expert_num = 4 - global_rank = 1 - redundancy_expert_num = 1 - topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - - redundancy_expert_ids = torch.tensor( - [ - 5, - ], - dtype=torch.int64, - device="cuda", - ) - redundancy_topk_ids_repair( - topk_ids=topk_ids, - redundancy_expert_ids=redundancy_expert_ids, - ep_expert_num=ep_expert_num, - global_rank=global_rank, - ) - - ans_topk_ids = torch.tensor( - [ - [0, 1, 2, 3], - [7, 9, 10, 11], - [1, 3, 5, 7], - ], - dtype=torch.int64, - device="cuda", - ) - ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids - new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids - ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( - (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 - ) - - assert torch.equal(topk_ids, ans_topk_ids) - - -def test_expert_id_counter(): - token_num = 256 - tok_ids = torch.randint( - low=0, - high=12, - size=(token_num, 8), - dtype=torch.int64, - device="cuda", - ) - expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) - - ans_expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - ids, counts = torch.unique(tok_ids.view(-1), return_counts=True) - ans_expert_counter[ids] = counts - - assert torch.equal(expert_counter, ans_expert_counter) - - # test speed - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - for _ in range(100): - tok_ids = torch.randint( - low=0, - high=12, - size=(token_num, 8), - dtype=torch.int64, - device="cuda", - ) - expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") - expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) - graph.replay() - - start_event = torch.cuda.Event(enable_timing=True) - start_event.record() - graph.replay() - end_event = torch.cuda.Event(enable_timing=True) - end_event.record() - torch.cuda.synchronize() - logger.info(f"expert_id_counter time cost: {start_event.elapsed_time(end_event)} ms") - - -if __name__ == "__main__": - pytest.main() diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py new file mode 100644 index 0000000000..bb3e97fb8e --- /dev/null +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -0,0 +1,3660 @@ +import builtins +import io +import threading +import time +from collections import deque +from contextlib import contextmanager +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, + build_logical_to_physical_map, + build_logical_to_physical_maps_for_layers, + _estimate_rank_load, + plan_redundant_experts, + select_improving_placements, +) +from lightllm.server.api_cli import make_argument_parser +from lightllm.server.core.objs.start_args_type import StartArgs +from lightllm.server.router.model_infer.infer_batch import g_infer_context +from lightllm.server.router.model_infer.mode_backend import ( + eplb_manager as manager_module, +) +from lightllm.server.router.model_infer.mode_backend import ( + eplb_transfer as transfer_module, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( + deepgemm_impl as deepgemm_module, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( + create_fuse_moe_impl, + FuseMoeMarlin, + FuseMoeTriton, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.base_impl import ( + FuseMoeBaseImpl, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( + fused_moe_weight as fused_weight_module, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + disable_eplb_model_init, + is_eplb_model_init_disabled, +) +from lightllm.common.eplb_utils import extract_eplb_expert_tensors +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + TransferStep, + _CudaBatchMemcpy, + _commit_staging_rows, + align_target_placement, + build_transfer_plan, +) + + +def _test_parallel_state( + *, + eplb=False, + num_logical_experts=128, + world_size=16, + num_redundant_experts_per_rank=1, + route_counter=None, + recording=False, + recorded_sample_count=0, +): + eplb_state = None + if eplb: + if route_counter is None: + route_counter = torch.zeros((2, num_logical_experts), dtype=torch.int64) + initial_layout_world_size = max(world_size, 2) + eplb_state = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( + num_logical_experts, + initial_layout_world_size, + num_redundant_experts_per_rank, + ), + logical_to_physical_map=torch.zeros((num_logical_experts, 2), dtype=torch.int32), + logical_replica_count=torch.ones(num_logical_experts, dtype=torch.int32), + route_counter=route_counter, + recording=recording, + recorded_sample_count=recorded_sample_count, + ) + return ExpertParallelState( + num_logical_experts=num_logical_experts, + world_size=world_size, + eplb=eplb_state, + ) + + +def _validated_expert_parallel_state( + *, + eplb=True, + n_routed_experts=4, + world_size=2, + num_redundant_experts_per_rank=1, + device="cpu", +): + runtime = None + if eplb: + runtime = EPLBState( + num_redundant_experts_per_rank=num_redundant_experts_per_rank, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( + n_routed_experts, + world_size, + num_redundant_experts_per_rank, + ), + logical_to_physical_map=torch.empty((n_routed_experts, 2), dtype=torch.int32, device=device), + logical_replica_count=torch.ones(n_routed_experts, dtype=torch.int32, device=device), + route_counter=torch.zeros((2, n_routed_experts), dtype=torch.int64, device=device), + ) + return ExpertParallelState( + num_logical_experts=n_routed_experts, + world_size=world_size, + eplb=runtime, + ) + + +def _set_expert_parallel_state(impl, state): + impl.expert_parallel_state = state + impl.eplb = state.eplb + impl._primary_weight_pack_cache = {} + + +def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): + """Reference the committed runtime logical-to-physical maps on CPU.""" + samples, layers, nodes, num_logical_experts = source_load.shape + ranks, redundant = placement.shape[1:] + num_experts_per_rank = num_logical_experts // ranks + num_physical_experts_per_rank = num_experts_per_rank + redundant + raw = torch.zeros((samples, layers, ranks, num_logical_experts), dtype=torch.float64) + for layer in range(layers): + for source_node in range(nodes): + logical_to_physical, replica_count = build_logical_to_physical_map( + placement[layer], + num_logical_experts, + source_rank=source_node * node_world_size, + node_world_size=node_world_size, + ) + for expert in range(num_logical_experts): + count = int(replica_count[expert].item()) + for physical_id in logical_to_physical[expert, :count].tolist(): + rank = physical_id // num_physical_experts_per_rank + raw[:, layer, rank, expert] += source_load[:, layer, source_node, expert] / count + return (torch.ceil(raw / alignment) * alignment).sum(dim=3) + + +def test_base_call_template_forwards_selection_and_capture_callback(): + class Impl(FuseMoeBaseImpl): + def _select_experts( + self, + input_tensor, + router_logits, + correction_bias, + top_k, + renormalize, + use_grouped_topk, + topk_group, + num_expert_group, + scoring_func, + per_expert_scale=None, + shared_expert_gate=None, + is_prefill=None, + preserve_logical_ids=False, + ): + seen["select"] = {"preserve_logical_ids": preserve_logical_ids} + return "weights", "physical_ids", "logical_ids" + + def _fused_experts( + self, + input_tensor, + w13, + w2, + topk_weights, + topk_ids, + router_logits=None, + is_prefill=None, + ): + seen["fused"] = {"topk_ids": topk_ids} + return "output" + + seen, captured = {}, [] + impl = Impl(4, 0, 1.0, SimpleNamespace()) + result = impl( + "input", + "logits", + "w13", + "w2", + None, + "softmax", + 2, + False, + False, + 0, + 0, + moe_capture_callback=captured.append, + ) + assert result == "output" + assert captured == ["logical_ids"] + assert seen["select"]["preserve_logical_ids"] + assert seen["fused"]["topk_ids"] == "physical_ids" + + +def test_parallel_state_derives_expert_layout(): + state = _validated_expert_parallel_state(eplb=True) + assert state.num_primary_experts_per_rank == 2 + assert state.num_total_physical_experts == 6 + + +def test_factory_selects_all_paths_and_requires_ep_state(monkeypatch): + plain_quant = SimpleNamespace(method_name="none") + marlin_quant = SimpleNamespace(method_name="awq_marlin") + monkeypatch.setattr(FuseMoeMarlin, "create_workspace", lambda self: None) + state = _validated_expert_parallel_state(eplb=False) + ep_impl = create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=plain_quant, + expert_parallel_state=state, + ) + assert isinstance(ep_impl, deepgemm_module.FuseMoeDeepGEMM) + assert ep_impl.expert_parallel_state is state + assert state.eplb is None + assert isinstance( + create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=plain_quant, + ), + FuseMoeTriton, + ) + assert isinstance( + create_fuse_moe_impl( + n_routed_experts=4, + num_fused_shared_experts=0, + routed_scaling_factor=1.0, + quant_method=marlin_quant, + ), + FuseMoeMarlin, + ) + + +def test_find_fused_moe_weights_discovers_direct_layer_attributes(monkeypatch): + class FakeFusedMoeWeight: + def __init__(self, layer_num, enable_ep_moe=True): + self.layer_num_ = layer_num + self.enable_ep_moe = enable_ep_moe + + monkeypatch.setattr(manager_module, "FusedMoeWeight", FakeFusedMoeWeight) + first = FakeFusedMoeWeight(3) + alternate = FakeFusedMoeWeight(1) + aliased = FakeFusedMoeWeight(2) + disabled = FakeFusedMoeWeight(0, enable_ep_moe=False) + model = SimpleNamespace( + trans_layers_weight=[ + SimpleNamespace(moe_weight=first), + SimpleNamespace(alternate_moe_weight=alternate), + SimpleNamespace(moe_weight=aliased, alternate_moe_weight=aliased), + SimpleNamespace(moe_weight=disabled), + ] + ) + + assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] + + +def test_eplb_redundant_experts_defaults_per_ep_rank(): + parser = make_argument_parser() + + assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 2 + assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 + assert StartArgs().eplb_num_redundant_experts_per_rank == 2 + + +@pytest.mark.parametrize( + ("num_logical_experts", "num_ranks", "num_redundant_experts_per_rank", "expected"), + [ + (8, 4, 2, [[2, 3], [4, 5], [6, 7], [0, 1]]), + (6, 3, 4, [[2, 3, 4, 5], [4, 5, 0, 1], [0, 1, 2, 3]]), + ], +) +def test_build_initial_redundant_expert_ids( + num_logical_experts, + num_ranks, + num_redundant_experts_per_rank, + expected, +): + actual = build_initial_redundant_expert_ids( + num_logical_experts, + num_ranks, + num_redundant_experts_per_rank, + ) + + assert actual.dtype == torch.int64 + assert actual.shape == (num_ranks, num_redundant_experts_per_rank) + assert torch.equal(actual, torch.tensor(expected, dtype=torch.int64)) + + +def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): + expert_load = torch.tensor( + [ + [100, 90, 80, 70, 60, 50, 40, 30], + [30, 40, 50, 60, 70, 80, 90, 100], + ] + ) + placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) + + for layer_placement in placement: + for rank, expert_ids in enumerate(layer_placement.tolist()): + assert len(expert_ids) == len(set(expert_ids)) + assert all(expert_id // 2 != rank for expert_id in expert_ids) + + +def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): + samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]) + placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) + candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] + + def critical(candidate): + return _estimate_rank_load(samples, candidate, expert_alignment=128).max(dim=2).values.sum() + + assert torch.equal(placement, torch.tensor([[[3], [0]]])) + assert critical(placement) == min(critical(candidate) for candidate in candidates) + + +def test_select_improving_placements_rejects_regressing_layer(): + expert_load = torch.tensor([[8649, 5740, 5002, 3441]]) + current = torch.tensor([[[2], [0]]]) + regressing_candidate = torch.tensor([[[1], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, regressing_candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, regressing_candidate).max() + / _estimate_rank_load(expert_load, regressing_candidate).mean() + ) + assert current_ratio.item() == pytest.approx(1.1007, abs=1e-4) + assert candidate_ratio.item() == pytest.approx(1.1184, abs=1e-4) + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): + expert_load = torch.tensor([[1, 2, 1, 17]]) + current = torch.tensor([[[3], [0]]]) + candidate = torch.tensor([[[3], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + assert current_ratio.item() == pytest.approx(1.047619, abs=1e-6) + assert candidate_ratio.item() == pytest.approx(1.0) + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): + expert_load = torch.tensor([[100, 129, 100, 129]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [1]]]) + + selected, improved, metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + + assert metrics["model_imbalance_ratio"] == pytest.approx(1.0) + assert metrics["candidate_rebalance_gain"] == pytest.approx(0.25) + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_insufficient_rebalance_gain(): + expert_load = torch.tensor([[1, 1, 6, 7]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + relative_improvement = (current_ratio - candidate_ratio) / current_ratio + assert current_ratio.item() == pytest.approx(1.4) + assert candidate_ratio.item() == pytest.approx(1.333333, abs=1e-6) + assert relative_improvement.item() == pytest.approx(0.047619, abs=1e-6) + assert not improved.item() + assert torch.equal(selected, current) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=0.04, + ) + + assert improved.item() + assert torch.equal(selected, candidate) + + +@pytest.mark.parametrize("rebalance_gain_threshold", [-0.01, 1.01, float("nan"), float("inf")]) +def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( + rebalance_gain_threshold, +): + expert_load = torch.tensor([[1, 1, 1, 2]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + with pytest.raises(ValueError, match="rebalance_gain_threshold"): + select_improving_placements( + expert_load, + current, + candidate, + rebalance_gain_threshold=rebalance_gain_threshold, + ) + + +def test_select_improving_placements_accepts_sufficient_rebalance_gain(): + expert_load = torch.tensor([[1, 1, 1, 2]]) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[3], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() + candidate_ratio = ( + _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() + ) + relative_improvement = (current_ratio - candidate_ratio) / current_ratio + assert current_ratio.item() == pytest.approx(1.2) + assert candidate_ratio.item() == pytest.approx(1.0) + assert relative_improvement.item() == pytest.approx(1 / 6) + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): + expert_load = torch.tensor([[1, 1, 1, 8]]) + current = torch.tensor([[[2], [0]]]) + raw_improving_candidate = torch.tensor([[[3], [0]]]) + + _, raw_improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, raw_improving_candidate, rebalance_gain_threshold=0.05 + ) + selected, aligned_improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, + current, + raw_improving_candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + + assert raw_improved.item() + assert not aligned_improved.item() + assert torch.equal(selected, current) + + +def test_estimate_rank_load_aligns_each_sample_before_accumulation(): + samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]) + placement = torch.tensor([[[2], [0]]]) + + per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) + accumulated = _estimate_rank_load(samples.sum(dim=0), placement, expert_alignment=128) + + assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) + assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) + assert torch.equal(accumulated[0], torch.tensor([128.0, 128.0])) + + +def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): + samples = torch.tensor( + [ + [[255, 220, 226, 254]], + [[172, 278, 51, 238]], + [[249, 291, 284, 183]], + ] + ) + current = torch.tensor([[[2], [0]]]) + mean_inflating_candidate = torch.tensor([[[2], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + samples, + current, + mean_inflating_candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + ) + current_load = _estimate_rank_load(samples, current, expert_alignment=128) + candidate_load = _estimate_rank_load(samples, mean_inflating_candidate, expert_alignment=128) + current_critical = current_load.max(dim=2).values.sum() + candidate_critical = candidate_load.max(dim=2).values.sum() + + assert candidate_load.mean(dim=2).sum() > current_load.mean(dim=2).sum() + assert current_critical == candidate_critical + assert not improved.item() + assert torch.equal(selected, current) + + +def test_select_improving_placements_accepts_five_percent_critical_reduction(): + samples = torch.tensor( + [ + [[13, 352, 348, 141]], + [[287, 175, 236, 179]], + [[316, 99, 266, 353]], + ] + ) + current = torch.tensor([[[2], [0]]]) + candidate = torch.tensor([[[2], [1]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + samples, current, candidate, rebalance_gain_threshold=0.05, expert_alignment=128 + ) + + assert improved.item() + assert torch.equal(selected, candidate) + + +def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): + # Layer 0 becomes better, but layer 1 dominates model critical load. The + # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]) + current = torch.tensor([[[2], [0]], [[2], [0]]]) + candidate = torch.tensor([[[3], [0]], [[2], [0]]]) + + selected, improved, metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + assert not torch.any(improved) + assert torch.equal(selected, current) + assert metrics["candidate_rebalance_gain"] == pytest.approx(0.5 / 13) + assert metrics["candidate_changed_layer_count"] == 1 + + +def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]) + current = torch.tensor([[[2], [0]], [[2], [0]]]) + candidate = torch.tensor([[[3], [0]], [[2], [0]]]) + + selected, improved, _metrics, _before_load, _after_load = select_improving_placements( + expert_load, current, candidate, rebalance_gain_threshold=0.05 + ) + + assert torch.equal(improved, torch.tensor([True, False])) + assert torch.equal(selected, candidate) + + +def test_logical_to_physical_map_has_at_most_one_slot_per_rank(): + redundant_expert_ids = torch.tensor([[2, 3], [0, 1]]) + logical_to_physical, replica_count = build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=4) + + assert logical_to_physical.shape == (4, 2) + assert torch.equal(replica_count, torch.full((4,), 2, dtype=torch.int64)) + + +def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count_without_source_rank(): + redundant_expert_ids = torch.tensor([[0], [1]]) + + with pytest.raises(AssertionError): + build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=5) + + +def test_logical_to_physical_map_prefers_source_node_replicas(): + # Four ranks, two ranks per node, two primary experts/rank and one + # redundant slot/rank. Expert 0 is primary on rank 0 and replicated on + # rank 2 (the other node). + redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) + rank0_map, rank0_count = build_logical_to_physical_map( + redundant, num_logical_experts=8, source_rank=0, node_world_size=2 + ) + rank1_map, rank1_count = build_logical_to_physical_map( + redundant, num_logical_experts=8, source_rank=1, node_world_size=2 + ) + fallback_redundant = torch.tensor([[4], [5], [0], [1], [2], [3]], dtype=torch.int64) + rank4_map, rank4_count = build_logical_to_physical_map(fallback_redundant, 12, source_rank=4, node_world_size=2) + + assert rank0_count[0].item() == rank1_count[0].item() == 1 + assert torch.equal(rank0_map[0, :1], torch.tensor([0])) + assert torch.equal(rank1_map[0, :1], torch.tensor([0])) + assert rank4_count[0].item() == 2 + assert set(rank4_map[0, :2].tolist()) == {0, 8} + + +def test_source_node_local_maps_fall_back_to_global_replicas(): + redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) + maps = [build_logical_to_physical_map(redundant, 8, source_rank=rank, node_world_size=2) for rank in range(4)] + assert maps[0][1][0].item() == maps[1][1][0].item() == 1 + assert maps[2][1][0].item() == maps[3][1][0].item() == 1 + assert maps[0][0][0, 0].item() == maps[1][0][0, 0].item() == 0 + assert maps[2][0][0, 0].item() == maps[3][0][0, 0].item() == 8 + + +def test_source_rank_rotates_selected_replica_order_without_changing_copies(): + # Expert 0 is primary on rank 0 and redundant on rank 1, so both ranks + # on node 0 have the same two local copies. Their source-rank phases + # must differ while their selected set/count remain identical. + redundant = torch.tensor([[1], [0], [3], [2]], dtype=torch.int64) + rank0_map, rank0_count = build_logical_to_physical_map(redundant, 4, source_rank=0, node_world_size=2) + rank1_map, rank1_count = build_logical_to_physical_map(redundant, 4, source_rank=1, node_world_size=2) + + assert rank0_count[0].item() == rank1_count[0].item() == 2 + assert set(rank0_map[0, :2].tolist()) == set(rank1_map[0, :2].tolist()) == {0, 3} + assert torch.equal(rank1_map[0, :2], torch.tensor([3, 0], dtype=torch.int32)) + + +@pytest.mark.parametrize( + "source_rank,node_world_size", + [(None, None), (0, 2), (1, 2), (2, 2), (3, 2)], +) +def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, node_world_size): + placements_by_layer = torch.tensor( + [ + [[4, 5], [0, 1], [0, 1], [2, 3]], + [[6, 7], [0, 1], [0, 1], [2, 3]], + [[4, 5], [0, 1], [0, 1], [2, 3]], + ], + dtype=torch.int64, + ) + + maps_by_layer, counts_by_layer = build_logical_to_physical_maps_for_layers( + placements_by_layer, + num_logical_experts=8, + source_rank=source_rank, + node_world_size=node_world_size, + ) + expected_by_layer = [ + build_logical_to_physical_map( + placement, + num_logical_experts=8, + source_rank=source_rank, + node_world_size=node_world_size, + ) + for placement in placements_by_layer + ] + + assert maps_by_layer.shape == (3, 8, 4) + assert counts_by_layer.shape == (3, 8) + assert maps_by_layer.dtype == counts_by_layer.dtype == torch.int32 + assert torch.equal(maps_by_layer, torch.stack([item[0] for item in expected_by_layer])) + assert torch.equal(counts_by_layer, torch.stack([item[1] for item in expected_by_layer])) + if source_rank is None: + # expert 0 的主副本在 rank 0,且 rank 1、2 都有一个冗余副本;第 1、2 + # 个冗余副本必须分别写入映射表的第 1、2 列,而不能互相覆盖。 + assert torch.equal(counts_by_layer[:, 0], torch.tensor([3, 3, 3], dtype=torch.int32)) + assert torch.equal( + maps_by_layer[:, 0, :3], + torch.tensor([[0, 6, 10], [0, 6, 10], [0, 6, 10]], dtype=torch.int32), + ) + positions = torch.arange(maps_by_layer.shape[-1]).view(1, 1, -1) + valid = positions < counts_by_layer.unsqueeze(-1) + assert torch.all(maps_by_layer[valid] >= 0) + assert torch.all(maps_by_layer[~valid] == -1) + + +def test_plan_redundant_experts_prefers_first_replica_on_new_node(): + # One redundant slot per rank leaves legal alternatives on both nodes; + # topology preference therefore puts every first replica away from its + # primary node before considering same-node duplicates. + placement = plan_redundant_experts( + torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]), + num_ranks=4, + num_redundant_experts_per_rank=1, + node_world_size=2, + ) + for rank, expert in enumerate(placement[0, :, 0].tolist()): + assert expert // 2 // 2 != rank // 2 + + +def test_plan_redundant_experts_single_node_matches_default_behavior(): + load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]) + default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) + single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) + assert torch.equal(single_node, default) + + +def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): + # Expert 0 has a copy on each node. Node 0 and node 1 issue unequal + # traffic, so collapsing them before planning produces the wrong result. + placement = torch.tensor([[[4], [5], [0], [1]]], dtype=torch.int64) + source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) + source_load[0, 0, 0, 0] = 256 + source_load[0, 0, 1, 0] = 128 + + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) + runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) + collapsed_global = _estimate_rank_load(source_load.sum(dim=2), placement, expert_alignment=128) + + assert torch.equal(predicted, runtime) + assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) + assert not torch.equal(predicted, collapsed_global) + + +def test_source_node_planner_constraints_and_real_critical_improvement(): + source_load = torch.tensor( + [ + [ + [ + [697, 451, 383, 536, 349, 404, 854, 425], + [861, 103, 166, 612, 444, 263, 910, 392], + ] + ], + [ + [ + [944, 457, 338, 108, 63, 525, 48, 216], + [439, 117, 837, 550, 833, 201, 729, 5], + ] + ], + [ + [ + [749, 159, 18, 723, 12, 700, 419, 51], + [112, 135, 8, 840, 40, 970, 90, 683], + ] + ], + ], + dtype=torch.int64, + ) + initial = build_initial_redundant_expert_ids(8, 4, 1).unsqueeze(0) + planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + + for rank, experts in enumerate(planned[0].tolist()): + assert len(experts) == len(set(experts)) == 1 + assert experts[0] // 2 != rank + + before = _estimate_rank_load(source_load, initial, expert_alignment=128, node_world_size=2) + after = _estimate_rank_load(source_load, planned, expert_alignment=128, node_world_size=2) + manual_before = _manual_runtime_rank_load(source_load, initial, 2, 128) + manual_after = _manual_runtime_rank_load(source_load, planned, 2, 128) + assert torch.equal(before, manual_before) + assert torch.equal(after, manual_after) + assert after.max(dim=2).values.sum() < before.max(dim=2).values.sum() + + +def test_source_node_select_uses_the_same_runtime_critical_prediction(): + source_load = torch.zeros((2, 1, 2, 8), dtype=torch.int64) + source_load[:, 0, 0, 0] = torch.tensor([1024, 768]) + source_load[:, 0, 1, 6] = torch.tensor([896, 1024]) + current = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) + candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) + selected, _improved, _metrics, _before_load, _after_load = select_improving_placements( + source_load, + current, + candidate, + rebalance_gain_threshold=0.05, + expert_alignment=128, + node_world_size=2, + ) + assert torch.equal( + _estimate_rank_load(source_load, selected, 128, 2), + _manual_runtime_rank_load(source_load, selected, 2, 128), + ) + + +def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: + """Count rank rows gaining an expert: one migrated row per new expert id.""" + moved = 0 + for layer in range(current.shape[0]): + for rank in range(current.shape[1]): + moved += len(set(target[layer, rank].tolist()) - set(current[layer, rank].tolist())) + return moved + + +def test_sticky_plan_reproduces_current_when_load_unchanged(): + generator = torch.Generator().manual_seed(7) + load = torch.randint(1, 1000, (3, 16, 32), generator=generator) + placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + + replanned = plan_redundant_experts( + load, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=placement, + stickiness=0.1, + ) + + assert torch.equal(replanned, placement) + for layer in range(placement.shape[0]): + assert build_transfer_plan(placement[layer], replanned[layer], 32, 4, 4) == [] + + +def test_sticky_plan_bounded_moves_under_small_perturbation(): + generator = torch.Generator().manual_seed(11) + load = torch.randint(100, 1000, (4, 16, 32), generator=generator) + placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + noise = torch.rand((4, 16, 32), generator=generator) * 0.1 + 0.95 + perturbed = (load.double() * noise).round().to(torch.int64) + + sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) + free = plan_redundant_experts(perturbed, 4, 2) + + sticky_moves = _count_moved_slots(placement, sticky) + free_moves = _count_moved_slots(placement, free) + assert sticky_moves <= placement.numel() // 4 + assert sticky_moves < free_moves + + def critical(candidate): + return _estimate_rank_load(perturbed, candidate).max(dim=2).values.sum() + + assert critical(sticky) <= critical(free) * 1.1 + + +def test_sticky_plan_still_churns_under_phase_shift(): + layers, experts = 8, 32 + before = torch.full((layers, experts), 10, dtype=torch.int64) + after = torch.full((layers, experts), 10, dtype=torch.int64) + offsets = torch.arange(4) + for layer in range(layers): + before[layer, (4 * layer + offsets) % experts] = 5000 + after[layer, (4 * layer + 16 + offsets) % experts] = 5000 + placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) + + replanned = plan_redundant_experts( + after, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=placement, + stickiness=0.1, + ) + + assert _count_moved_slots(placement, replanned) > placement.numel() // 2 + + +def test_transfer_plan_slot_permutation_is_free(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) + + assert torch.equal(align_target_placement(current, target), current) + assert build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) == [] + + +def test_align_target_placement_keeps_retained_experts_in_live_slots(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) + + canonical = align_target_placement(current, target) + + assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) + + +def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): + num_logical_experts = 8 + world_size = 4 + num_redundant_slots_per_rank = 2 + num_experts_per_rank = num_logical_experts // world_size + num_physical_experts_per_rank = num_experts_per_rank + num_redundant_slots_per_rank + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) + canonical = align_target_placement(current, target) + plan = build_transfer_plan(current, canonical, num_logical_experts, world_size, node_world_size=2) + + # Label every current physical row by its resident logical expert, then + # apply the transfer plan from a frozen source snapshot just as staging + # copies do before the destination rows are published. + source_rows = [ + list(range(rank * num_experts_per_rank, (rank + 1) * num_experts_per_rank)) + current[rank].tolist() + for rank in range(world_size) + ] + live_rows = [row.copy() for row in source_rows] + for step in plan: + live_rows[step.dst_rank][num_experts_per_rank + step.dst_slot] = source_rows[step.src_rank][step.src_local_row] + + logical_to_physical, replica_count = build_logical_to_physical_map(canonical, num_logical_experts) + for logical_expert, count in enumerate(replica_count.tolist()): + for physical_id in logical_to_physical[logical_expert, :count].tolist(): + rank, row = divmod(physical_id, num_physical_experts_per_rank) + assert live_rows[rank][row] == logical_expert + + +def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.world_size = 4 + manager.node_world_size = 2 + manager.num_logical_experts = 8 + manager.num_redundant_experts_per_rank = 2 + manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) + manager.placement_stickiness = 0.1 + manager.rebalance_gain_threshold = 0.05 + manager.evaluation_group = object() + candidate = torch.tensor([[[5, 4], [7, 6], [1, 0], [3, 2]]]) + broadcasts = [] + + def fixed_selector(*_args, **_kwargs): + rank_load = torch.full((1, 4), 100.0) + return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load + + def record_broadcast(result_list, **_kwargs): + broadcasts.append(result_list[0]) + + monkeypatch.setattr( + manager_module, + "plan_redundant_experts", + lambda *_args, **_kwargs: candidate.clone(), + ) + monkeypatch.setattr(manager_module, "select_improving_placements", fixed_selector) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) + + result = manager._plan_and_broadcast(torch.full((1, 1, 2, 8), 100, dtype=torch.int64)) + + assert torch.equal(result["placement"], manager.current_placement) + assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) + + +def test_stickiness_zero_matches_legacy(): + generator = torch.Generator().manual_seed(17) + load = torch.randint(1, 1000, (2, 8, 16), generator=generator) + legacy = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() + + replanned = plan_redundant_experts( + load, + num_ranks=4, + num_redundant_experts_per_rank=2, + current_placement=unrelated, + stickiness=0.0, + ) + + assert torch.equal(replanned, legacy) + + +def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast( + monkeypatch, +): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.global_rank = 0 + manager.world_size = 2 + manager.node_world_size = 2 + manager.num_logical_experts = 8 + manager.num_redundant_experts_per_rank = 2 + manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) + manager.placement_stickiness = 0.1 + manager.rebalance_gain_threshold = 0.05 + manager.evaluation_group = object() + broadcasted = [] + + def broken_planner(*_args, **_kwargs): + raise RuntimeError("planner boom") + + def record_broadcast(result_list, **_kwargs): + broadcasted.append(result_list[0]) + + monkeypatch.setattr(manager_module, "plan_redundant_experts", broken_planner) + monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) + + with pytest.raises(RuntimeError, match="EPLB planner failed on rank zero") as exc_info: + manager._plan_and_broadcast(torch.full((1, 1, 1, 8), 100, dtype=torch.int64)) + + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "planner boom" + assert broadcasted == [{"kind": "error", "message": "RuntimeError: planner boom"}] + + +def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_evaluates_step_twenty( + monkeypatch, +): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=1, + num_logical_experts=4, + world_size=1, + ), + }, + )() + ] + manager.in_flight = False + manager.prefill_steps = 15 + manager.step_interval = 20 + manager.sampling_interval = manager.step_interval + manager.evaluation_group = object() + manager.num_logical_experts = 4 + manager.global_rank = 1 + manager.evaluation_in_flight = False + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + + recordings, resets, started = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + monkeypatch.setattr( + manager_module.torch.cuda, + "synchronize", + lambda: pytest.fail("step must not synchronize CUDA"), + ) + + manager.step() + assert manager.prefill_steps == 16 + assert recordings == [True] + assert resets == [True] + assert manager._steady_collection_end_step == 20 + + for _ in range(3): + manager.step() + assert manager.prefill_steps == 19 + assert recordings == [True] + assert started == [] + + manager.step() + assert manager.prefill_steps == 20 + assert manager._steady_collection_end_step is None + assert started == [True] + + +def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary( + monkeypatch, +): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = False + manager.prefill_steps = 0 + manager.step_interval = 20 + manager.sampling_interval = 3 + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + recordings, resets, started = [], [], [] + manager._set_recording = lambda enabled: recordings.append((manager.prefill_steps, enabled)) + manager._reset_recorded_samples = lambda: resets.append(manager.prefill_steps) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(manager.prefill_steps)) + + manager._prepare_next_sampling_window() + assert resets == [0] + assert recordings == [(0, True)] + assert manager._steady_collection_end_step == 3 + + manager.step() + manager.step() + assert started == [] + manager.step() + assert started == [3] + + +def test_eplb_step_does_not_start_a_second_evaluation_while_one_is_pending(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=1, + num_logical_experts=4, + world_size=1, + ), + }, + )() + ] + manager.in_flight = False + manager.prefill_steps = 1 + manager.step_interval = 2 + manager.sampling_interval = manager.step_interval + manager.evaluation_group = object() + manager.num_logical_experts = 4 + manager.global_rank = 1 + manager.evaluation_in_flight = True + + started = [] + monkeypatch.setattr(manager, "_poll_evaluation", lambda: True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + + manager.step() + + assert manager.prefill_steps == 1 + assert started == [] + + +def test_evaluation_no_improvement_logs_model_fields_without_reopening_interval_window( + monkeypatch, +): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 0 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.weights = [] + manager._eplb_states = [] + recordings, logs = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) + + assert not manager._poll_evaluation() + assert recordings == [False] + assert "model_imbalance_ratio" in logs[0][0] + assert "candidate_rebalance_gain" in logs[0][0] + assert "candidate_changed_layer_count" in logs[0][0] + assert "actual_changed_layer_count" in logs[0][0] + assert "next_sampling_interval" in logs[0][0] + assert manager.sampling_interval == 80 + + +def test_interval_one_rearms_after_evaluation_but_never_evaluates_empty_counter( + monkeypatch, +): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.prefill_steps = 1 + manager.step_interval = 1 + manager.sampling_interval = 1 + manager._steady_collection_end_step = None + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 1 + manager.weights = [] + manager._eplb_states = [] + recordings, starts = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._start_evaluation = lambda: starts.append(True) + manager._evaluation_ready_on_all_ranks = lambda: True + + manager.poll() + + # The no-improvement backoff changes interval 1 to 4. The clamped + # steady window arms immediately but still waits for boundary step 5. + assert recordings == [True] + assert manager._steady_collection_end_step is not None + assert manager._steady_collection_end_step == 5 + assert starts == [] + assert manager.prefill_steps == 1 + + for _ in range(3): + manager.step() + assert starts == [] + manager.step() + assert starts == [True] + + +def test_evaluation_worker_error_is_raised_by_main_thread(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = RuntimeError("planner failed") + manager._evaluation_result = None + manager._evaluation_thread = DoneThread() + + with pytest.raises(RuntimeError, match="planner failed"): + manager._poll_evaluation() + + +def test_evaluation_state_is_cleared_before_second_round(monkeypatch): + class DoneThread: + def join(self): + pass + + class PendingThread: + def __init__(self, **_kwargs): + self.started = False + + def start(self): + self.started = True + + def join(self): + pytest.fail("pending worker must not be joined") + + class Event: + def record(self, _stream): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 1, + } + manager.global_rank = 1 + manager.prefill_steps = 0 + manager.step_interval = 1 + manager.sampling_interval = 1 + manager.weights = [] + manager._eplb_states = [] + manager._set_recording = lambda _enabled: None + monkeypatch.setattr(manager_module.threading, "Thread", PendingThread) + monkeypatch.setattr(manager_module.torch.cuda, "Event", Event) + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) + + assert not manager._poll_evaluation() + assert manager._evaluation_result is None + assert manager._evaluation_error is None + + manager._start_evaluation() + assert manager.evaluation_in_flight + assert manager._evaluation_result is None + assert manager._poll_evaluation() # New worker has not produced a result. + + +def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch): + counters = [ + torch.tensor([[10, 11], [20, 21], [30, 31]], dtype=torch.int64), + torch.tensor([[40, 41], [50, 51], [60, 61]], dtype=torch.int64), + ] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=5, + num_logical_experts=2, + world_size=1, + ), + }, + )() + for counter in counters + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.evaluation_group = object() + metadata_sizes = [] + + def all_reduce(metadata, **_kwargs): + metadata_sizes.append(metadata.numel()) + + monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) + + samples = manager._collect_local_samples() + + assert metadata_sizes == [4] + assert torch.equal( + samples, + torch.tensor( + [ + [[30, 31], [60, 61]], + [[10, 11], [40, 41]], + [[20, 21], [50, 51]], + ], + dtype=torch.int64, + ), + ) + + +def test_manager_collects_only_two_recent_sparse_samples(monkeypatch): + counters = [ + torch.tensor([[10, 11], [20, 21], [30, 31], [40, 41]], dtype=torch.int64), + torch.tensor([[50, 51], [60, 61], [70, 71], [80, 81]], dtype=torch.int64), + ] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=2, + num_logical_experts=2, + world_size=1, + ), + }, + )() + for counter in counters + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.evaluation_group = object() + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **kwargs: None) + + samples = manager._collect_local_samples() + + assert torch.equal( + samples, + torch.tensor([[[10, 11], [50, 51]], [[20, 21], [60, 61]]], dtype=torch.int64), + ) + + +def test_eplb_counter_capacity_covers_default_dense_interval(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + monkeypatch.setattr(fused_weight_module, "get_prefill_eplb_step_interval", lambda: 20) + monkeypatch.setattr(fused_weight_module, "get_node_world_size", lambda: 2) + monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) + original_zeros = torch.zeros + + def cpu_zeros(*shape, **kwargs): + kwargs.pop("device", None) + return original_zeros(*shape, **kwargs) + + monkeypatch.setattr(fused_weight_module.torch, "zeros", cpu_zeros) + + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state.eplb.route_counter.shape == (40, 4) + + +def test_steady_sampling_resets_fixed_ring_without_retained_history(): + counter = torch.ones((8, 4), dtype=torch.int64) + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + state = _test_parallel_state( + eplb=True, + route_counter=counter, + recording=False, + recorded_sample_count=99, + num_logical_experts=4, + world_size=1, + ).eplb + manager._eplb_states = [state] + + manager._reset_recorded_samples() + manager._reset_recorded_samples() + + assert state.route_counter.shape == (8, 4) + assert torch.count_nonzero(state.route_counter) == 0 + assert state.recorded_sample_count == 0 + assert not hasattr(manager, "_retained_local_samples") + assert not hasattr(manager, "_sample_history") + + +def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": False, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state is not None + assert weight.expert_parallel_state.eplb is None + assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts + assert weight._initial_redundant_expert_ids == [] + assert not hasattr(weight, "route_counter") + assert not hasattr(weight, "routed_expert_counter_tensor") + + +def test_disable_eplb_model_init_skips_eplb_state(monkeypatch): + args = type( + "Args", + (), + {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, + )() + weight = object.__new__(fused_weight_module.FusedMoeWeight) + weight.n_routed_experts = 4 + weight.global_world_size = 2 + weight.global_rank_ = 0 + weight.enable_ep_moe = True + monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) + monkeypatch.setattr( + fused_weight_module, + "build_initial_redundant_expert_ids", + lambda *args, **kwargs: pytest.fail("disabled scope must not initialize EPLB"), + ) + + with disable_eplb_model_init(): + weight._init_expert_parallel_state() + + assert weight.expert_parallel_state.eplb is None + assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts + assert weight._initial_redundant_expert_ids == [] + + +def test_disable_eplb_model_init_scope_restores_after_exception(): + assert not is_eplb_model_init_disabled() + with disable_eplb_model_init(): + assert is_eplb_model_init_disabled() + with disable_eplb_model_init(): + assert is_eplb_model_init_disabled() + assert is_eplb_model_init_disabled() + + with pytest.raises(RuntimeError): + with disable_eplb_model_init(): + raise RuntimeError + assert not is_eplb_model_init_disabled() + + +def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 4), dtype=torch.int64), + num_logical_experts=4, + world_size=1, + ) + }, + )() + ] + manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager.global_rank = 2 + manager.world_size = 4 + manager.node_world_size = 2 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 4 + manager.num_redundant_experts_per_rank = 1 + manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0) + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + local = torch.full((1, 1, 4), 100, dtype=torch.int64) + manager._collect_local_samples = lambda: local + seen = {} + + def all_reduce(tensor, **kwargs): + seen["before"] = tensor.clone() + seen["group"] = kwargs["group"] + # Simulate source node 0's contribution from the other ranks. + tensor[:, :, 0].fill_(100) + + monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + + def plan_and_broadcast(global_load): + seen["global_load"] = global_load.clone() + return {"kind": "insufficient"} + + manager._plan_and_broadcast = plan_and_broadcast + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert seen["group"] is manager.evaluation_group + expected_local = local + assert seen["before"].shape == (1, 1, 2, 4) + assert torch.equal(seen["before"][:, :, 0], torch.zeros_like(expected_local)) + assert torch.equal(seen["before"][:, :, 1], expected_local) + assert torch.equal(seen["global_load"][:, :, 0], torch.full_like(expected_local, 100)) + assert torch.equal(seen["global_load"][:, :, 1], expected_local) + assert manager._evaluation_error is None + assert manager._evaluation_result["recorded_sample_count"] == 1 + assert manager._evaluation_result["sample_window_steps"] == 4 + + +def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 4), dtype=torch.int64), + num_logical_experts=4, + world_size=4, + ) + }, + )() + for _ in range(3) + ] + manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] + manager.global_rank = 1 + manager.world_size = 4 + manager.node_world_size = 2 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 4 + manager.num_redundant_experts_per_rank = 1 + manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + prepare_calls = [] + prepared_batches = object() + + def prepare_transfer(layer_plans): + prepare_calls.append(layer_plans) + return prepared_batches + + manager.transfer = SimpleNamespace(prepare_transfer=prepare_transfer) + manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) + planned_placement = torch.tensor( + [ + [[3], [0], [1], [2]], + [[2], [3], [0], [1]], + [[1], [2], [3], [0]], + ], + dtype=torch.int64, + ) + manager._plan_and_broadcast = lambda _global_load: { + "kind": "planned", + "placement": planned_placement, + "improved": torch.tensor([True, False, True]), + } + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + calls = [] + original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers + + def build_maps_for_layers(*args, **kwargs): + calls.append(args[0].shape) + return original_build_maps_for_layers(*args, **kwargs) + + monkeypatch.setattr( + manager_module, + "build_logical_to_physical_maps_for_layers", + build_maps_for_layers, + ) + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert manager._evaluation_error is None + assert calls == [torch.Size([2, 4, 1])] + metadata = manager._evaluation_result["metadata"] + assert metadata[1] is None + assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] + assert prepare_calls == [manager._evaluation_result["layer_plans"]] + assert manager._evaluation_result["prepared_batches"] is prepared_batches + for layer_index in (0, 2): + item = metadata[layer_index] + expected = build_logical_to_physical_map( + planned_placement[layer_index], + 4, + source_rank=manager.global_rank, + node_world_size=manager.node_world_size, + ) + assert torch.equal(item[0], expected[0]) + assert torch.equal(item[1], expected[1]) + + +def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.weights = [ + type( + "Weight", + (), + { + "expert_parallel_state": _test_parallel_state( + eplb=True, + route_counter=torch.zeros((1, 2), dtype=torch.int64), + num_logical_experts=2, + world_size=1, + ) + }, + )() + ] + manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] + manager.global_rank = 0 + manager.world_size = 1 + manager.node_world_size = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.num_logical_experts = 2 + manager.current_placement = torch.tensor([[[0], [1]]], dtype=torch.int64) + manager.evaluation_group = object() + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager._continuous_collection_end_step = None + manager._collect_local_samples = lambda: torch.ones((1, 1, 2), dtype=torch.int64) + manager._plan_and_broadcast = lambda _global_load: { + "kind": "planned", + "placement": torch.tensor([[[0], [1]]], dtype=torch.int64), + "improved": torch.tensor([True]), + } + + def fail_prepare(_layer_plans): + raise RuntimeError("prepare failed") + + manager.transfer = SimpleNamespace(prepare_transfer=fail_prepare) + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) + monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) + monkeypatch.setattr(manager_module, "build_transfer_plan", lambda *_args, **_kwargs: []) + + manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) + + assert manager._evaluation_result is None + assert isinstance(manager._evaluation_error, RuntimeError) + assert str(manager._evaluation_error) == "prepare failed" + + +def test_decode_dispatch_keeps_logical_ids_and_uses_logical_expert_count(monkeypatch): + class Buffer: + def low_latency_dispatch(self, **kwargs): + calls.append(kwargs) + return "recv", "masked", "handle", "event", "hook" + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.quant_method = type("Quant", (), {"method_name": "fp8"})() + impl.n_routed_experts = 128 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + impl._select_experts = lambda **_kwargs: ( + torch.ones((1, 2)), + torch.tensor([[0, 127]], dtype=torch.int32), + torch.tensor([[0, 127]], dtype=torch.int32), + ) + calls = [] + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_decode", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_low_latency_buffer", Buffer()) + + result = impl.low_latency_dispatch( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + False, + 2, + False, + 0, + 0, + "softmax", + ) + + assert result[2].tolist() == [[0, 127]] + assert calls[0]["num_experts"] == 128 + + +def test_decode_select_does_not_clone_or_map_eplb_topk_ids(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + topk_ids = torch.tensor([[3, 127]], dtype=torch.int32) + monkeypatch.setattr(topk_select, "select_experts", lambda **_kwargs: (torch.ones((1, 2)), topk_ids)) + _, selected, origin = impl._select_experts( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + 2, + False, + False, + 0, + 0, + "softmax", + is_prefill=False, + ) + + assert selected.data_ptr() == origin.data_ptr() + + +def test_eplb_prefill_uses_single_fused_path_for_global_topk(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + impl.quant_method = object() + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + physical_ids = torch.tensor([[130, 131]], dtype=torch.long) + calls = [] + + def fused_topk(**kwargs): + calls.append(kwargs) + return torch.ones((1, 2)), physical_ids, None + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") + + weights, topk_idx, qinput = impl.select_experts_and_quant_input( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + object(), + False, + 2, + False, + 0, + 0, + "softmax", + ) + + assert weights.tolist() == [[1.0, 1.0]] + assert topk_idx is physical_ids + assert topk_idx.dtype is torch.long + assert qinput == "qinput" + assert not calls[0]["use_grouped_topk"] + assert not calls[0]["return_logical_ids"] + + +def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + class Buffer: + def dispatch(self, _qinput, **kwargs): + calls.append(kwargs) + return ( + (torch.empty((4, 2)),), + "recv_idx", + "recv_weight", + SimpleNamespace(num_recv_tokens_per_expert_list=[4]), + SimpleNamespace(current_stream_wait=lambda: None), + ) + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + impl.quant_method = object() + state = _test_parallel_state( + eplb=True, + route_counter=torch.zeros((3, 128), dtype=torch.int64), + recording=True, + ) + _set_expert_parallel_state(impl, state) + impl.ep_balance_counters = None + calls, fused_calls = [], [] + physical_ids = torch.tensor([[130, 131]], dtype=torch.long) + + def fused_topk(**kwargs): + fused_calls.append(kwargs) + return torch.ones((1, 2)), physical_ids, None + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_prefill", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) + + weights, topk_idx, qinput = impl.select_experts_and_quant_input( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + object(), + True, + 2, + False, + 1, + 8, + "sigmoid", + ) + caller_event = object() + impl.dispatch( + qinput, + topk_idx, + weights, + overlap_event=caller_event, + ) + + assert topk_idx is physical_ids + assert len(fused_calls) == 1 + assert fused_calls[0]["sample_index"] == 0 + assert fused_calls[0]["record_load"] + assert state.eplb.recorded_sample_count == 1 + assert calls[0]["topk_idx"] is physical_ids + assert calls[0]["topk_idx"].dtype is torch.long + assert calls[0]["previous_event"] is caller_event + + +def test_prefill_dispatch_preserves_event(monkeypatch): + class Buffer: + def dispatch(self, _qinput, **kwargs): + calls.append(kwargs) + return ( + (torch.empty((4, 2)),), + "recv_idx", + "recv_weight", + SimpleNamespace(num_recv_tokens_per_expert_list=[4]), + SimpleNamespace(current_stream_wait=lambda: None), + ) + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + impl.ep_balance_counters = None + calls = [] + caller_event = object() + monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) + monkeypatch.setattr( + deepgemm_module, + "get_deepep_num_max_dispatch_tokens_per_rank_prefill", + lambda: 16, + ) + monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) + + impl.dispatch( + "qinput", + torch.tensor([[1, 2]], dtype=torch.long), + torch.ones((1, 2)), + caller_event, + ) + + assert calls[0]["previous_event"] is caller_event + assert calls[0]["topk_idx"].dtype is torch.long + + +def test_deepgemm_constructor_configures_eplb(): + state = _validated_expert_parallel_state() + impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), expert_parallel_state=state) + assert impl.expert_parallel_state is state + + +def test_prefill_eplb_returns_requested_logical_ids(monkeypatch): + from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk + + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.routed_scaling_factor = 1.0 + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) + + def fused_topk(**kwargs): + assert kwargs["return_logical_ids"] + return ( + torch.ones((1, 2)), + torch.tensor([[13, 14]], dtype=torch.int32), + torch.tensor([[3, 4]], dtype=torch.int32), + ) + + monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) + _, physical_ids, logical_ids = impl._select_experts( + torch.empty((1, 4)), + torch.empty((1, 128)), + None, + 2, + False, + False, + 0, + 0, + "softmax", + is_prefill=True, + preserve_logical_ids=True, + ) + + assert physical_ids.tolist() == [[13, 14]] + assert logical_ids.tolist() == [[3, 4]] + + +def test_decode_masked_group_gemm_uses_primary_rows_only_when_eplb_is_enabled( + monkeypatch, +): + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + _set_expert_parallel_state(impl, _test_parallel_state(eplb=True, num_logical_experts=8, world_size=1)) + captured = {} + + def masked(*args, **kwargs): + captured["w13"] = args[3] + captured["w13_scale"] = args[4] + captured["w2"] = args[5] + captured["w2_scale"] = args[6] + return "out" + + monkeypatch.setattr(deepgemm_module, "masked_group_gemm", masked) + pack = lambda: type( + "Pack", + (), + {"weight": torch.empty((10, 4)), "weight_scale": torch.empty((10, 1))}, + )() + + assert impl.masked_group_gemm((torch.empty((1, 4)),), pack(), pack(), torch.empty(8), torch.float16, 1) == "out" + assert captured["w13"].shape[0] == captured["w2"].shape[0] == 8 + assert captured["w13_scale"].shape[0] == captured["w2_scale"].shape[0] == 8 + + +def test_decode_fused_experts_uses_cached_primary_weight_packs_and_logical_experts( + monkeypatch, +): + impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) + impl.n_routed_experts = 128 + _set_expert_parallel_state( + impl, + _test_parallel_state(eplb=True, num_logical_experts=128, world_size=16, num_redundant_experts_per_rank=2), + ) + impl.quant_method = object() + impl.ep_balance_counters = None + captured = [] + + def fused(**kwargs): + captured.append(kwargs) + return "out" + + monkeypatch.setattr(deepgemm_module, "fused_experts", fused) + pack = lambda: type( + "Pack", + (), + { + "weight": torch.empty((10, 4)), + "weight_scale": torch.empty((10, 1)), + "weight_zero_point": None, + }, + )() + w13, w2 = pack(), pack() + + for _ in range(2): + assert ( + impl._fused_experts( + torch.empty((1, 4)), + w13, + w2, + torch.ones((1, 2)), + torch.zeros((1, 2), dtype=torch.int64), + is_prefill=False, + ) + == "out" + ) + + assert [call["num_experts"] for call in captured] == [128, 128] + assert all(call["w13"].weight.shape[0] == call["w2"].weight.shape[0] == 8 for call in captured) + assert captured[0]["w13"] is captured[1]["w13"] + assert captured[0]["w2"] is captured[1]["w2"] + + +def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): + current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) + target = current.clone() + target[0, 0] = 6 # primary r3, but r1 replica is on r0's node. + target[2, 1] = 4 # primary r2 is local to destination r2. + plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) + by_dst = {(step.dst_rank, step.dst_slot): step for step in plan} + assert len(by_dst) == 2 + assert by_dst[0, 0] == TransferStep(0, 0, 1, 2) + assert by_dst[2, 1] == TransferStep(2, 1, 2, 0) + + +def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): + current = torch.tensor([[0, 1], [2, 3], [4, 5], [4, 7]]) + target = current.clone() + target[0, 0] = 4 + target[0, 1] = 4 + first = build_transfer_plan(current, target, 8, 4, 2) + second = build_transfer_plan(current, target, 8, 4, 2) + assert first == second + selected = [step for step in first if step.dst_rank == 0] + assert [(step.src_rank, step.src_local_row) for step in selected] == [ + (2, 0), + (3, 2), + ] + + +def test_extract_expert_tensors_includes_weight_scale_and_zero_point_in_order(): + class Pack: + def __init__(self, offset, scale=True, zero=True): + self.weight = torch.full((3, 2), offset) + self.weight_scale = torch.full((3, 1), offset + 1) if scale else None + self.weight_zero_point = torch.full((3, 1), offset + 2) if zero else None + + weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero=False)})() + tensors = extract_eplb_expert_tensors(weight) + assert [name for name, _ in tensors] == [ + "w13.weight", + "w13.weight_scale", + "w13.weight_zero_point", + "w2.weight", + ] + + +def test_commit_staging_rows_only_overwrites_redundant_rows(): + live = torch.arange(20).reshape(5, 4) + staging = torch.full((2, 4), -1) + _commit_staging_rows( + live, + staging, + num_experts_per_rank=3, + changed_dst_slots=(0, 1), + ) + assert torch.equal(live[:3], torch.arange(12).reshape(3, 4)) + assert torch.equal(live[3:], staging) + + +def test_commit_staging_rows_preserves_unchanged_destination_slots(): + live = torch.arange(28).reshape(7, 4) + staging = torch.tensor([[-1, -1, -1, -1], [-2, -2, -2, -2], [-3, -3, -3, -3], [-4, -4, -4, -4]]) + original = live.clone() + + _commit_staging_rows( + live, + staging, + num_experts_per_rank=3, + changed_dst_slots=(3, 1), + ) + + assert torch.equal(live[:4], original[:4]) + assert torch.equal(live[4], staging[1]) + assert torch.equal(live[5], original[5]) + assert torch.equal(live[6], staging[3]) + + +def test_commit_staging_rows_merges_contiguous_changed_slots(): + copies = [] + + class View: + def __init__(self, owner, start, length): + self.owner = owner + self.start = start + self.length = length + + def copy_(self, source, **_kwargs): + copies.append( + ( + self.owner, + self.start, + self.length, + source.owner, + source.start, + source.length, + ) + ) + + class Tensor: + def __init__(self, owner, rows): + self.owner = owner + self.shape = (rows,) + + def narrow(self, _dim, start, length): + return View(self.owner, start, length) + + _commit_staging_rows( + Tensor("live", 20), + Tensor("staging", 4), + num_experts_per_rank=10, + changed_dst_slots=(3, 1, 2), + ) + + assert copies == [("live", 11, 3, "staging", 1, 3)] + + +def test_manager_inflight_ready_gate_commits_ordered_prefix_and_propagates_worker_error( + monkeypatch, +): + class Transfer: + def __init__(self): + self.pending = [(0, 0), (1, 1), (2, 2)] + self.commits = [] + self.finished = 0 + + def pending_layers(self): + return self.pending + + def commit(self, layer, buffer_index, post_copy=None): + assert self.pending.pop(0) == (layer, buffer_index) + self.commits.append((layer, buffer_index)) + if post_copy is not None: + post_copy() + + def finish(self): + self.finished += 1 + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer() + manager.in_flight = True + manager.world_size = 2 + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0, 1, 2] + manager._commit_layer_metadata = lambda layer: committed.append(layer) + manager._finish_rebalance = lambda: finished.append(True) + committed, finished = [], [] + operations = [] + current_stream_calls = [] + + class CurrentStream: + def wait_stream(self, stream): + operations.append(("wait", stream)) + + overlap_stream = object() + + def current_stream(): + current_stream_calls.append(True) + return CurrentStream() + + monkeypatch.setattr(manager_module.torch.cuda, "current_stream", current_stream) + monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) + + original_commit = manager.transfer.commit + + def record_commit(*args, **kwargs): + operations.append(("commit", args[0])) + return original_commit(*args, **kwargs) + + manager.transfer.commit = record_commit + + def set_global_ready(count): + return lambda tensor, **kwargs: tensor.fill_(count) + + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) + manager._poll_in_flight() + assert manager.transfer.commits == [] + assert operations == [] + assert current_stream_calls == [] + + # Local rank has three prefetched layers, but global MIN-ready only permits two. + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(2)) + manager._poll_in_flight() + assert manager.transfer.commits == [(0, 0), (1, 1)] + assert committed == [0, 1] + assert not finished + assert manager.transfer.finished == 0 + assert operations == [("wait", overlap_stream), ("commit", 0), ("commit", 1)] + assert current_stream_calls == [True] + + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + manager._poll_in_flight() + assert finished == [True] + assert manager.transfer.finished == 1 + assert operations == [ + ("wait", overlap_stream), + ("commit", 0), + ("commit", 1), + ("wait", overlap_stream), + ("commit", 2), + ] + assert current_stream_calls == [True, True] + + manager.in_flight_layers = [3] + manager.transfer.pending = [(9, 0)] + monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) + with pytest.raises(RuntimeError, match="does not match expected"): + manager._poll_in_flight() + + class BrokenTransfer: + def pending_layers(self): + raise RuntimeError("boom") + + manager.transfer = BrokenTransfer() + encoded_statuses = [] + + def retain_local_error(tensor, **_kwargs): + encoded_statuses.append(int(tensor.item())) + + monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) + with pytest.raises(RuntimeError, match="EPLB transfer worker failed on this rank") as exc_info: + manager._poll_in_flight() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "boom" + assert encoded_statuses == [manager_module.EPLB_CONTROL_ERROR] + + +def test_manager_inflight_remote_worker_error_does_not_commit(monkeypatch): + class Transfer: + def __init__(self): + self.commits = [] + + def pending_layers(self): + return [(0, 0)] + + def commit(self, *args): + self.commits.append(args) + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer() + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0] + manager._commit_layer_metadata = lambda _layer: None + manager._finish_rebalance = lambda: None + statuses = [] + + def remote_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + tensor.fill_(manager_module.EPLB_CONTROL_ERROR) + + monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) + + with pytest.raises(RuntimeError, match="EPLB transfer worker failed on another rank"): + manager._poll_in_flight() + assert statuses == [1] + assert manager.transfer.commits == [] + + +def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager._evaluation_lock = threading.Lock() + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager._evaluation_error = RuntimeError("evaluation boom") + manager._evaluation_result = None + statuses = [] + + def retain_local_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + + monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) + with pytest.raises(RuntimeError, match="EPLB evaluation failed on this rank") as exc_info: + manager._evaluation_ready_on_all_ranks() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "evaluation boom" + assert statuses == [manager_module.EPLB_CONTROL_ERROR] + + manager._evaluation_error = None + manager._evaluation_result = {"kind": "no_improvement"} + statuses.clear() + + def remote_error(tensor, **_kwargs): + statuses.append(int(tensor.item())) + tensor.fill_(manager_module.EPLB_CONTROL_ERROR) + + monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) + with pytest.raises(RuntimeError, match="EPLB evaluation failed on another rank"): + manager._evaluation_ready_on_all_ranks() + assert statuses == [1] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_manager_inflight_commit_orders_live_weights_between_overlap_forwards( + monkeypatch, +): + class Transfer: + def __init__(self, live, staging): + self.live = live + self.staging = staging + self.pending = [(0, 0)] + + def pending_layers(self): + return self.pending + + def commit(self, layer, buffer_index, post_copy=None): + assert self.pending.pop(0) == (layer, buffer_index) + self.live.copy_(self.staging, non_blocking=True) + if post_copy is not None: + post_copy() + + def finish(self): + pass + + live = torch.tensor([1.0], device="cuda") + staging = torch.tensor([2.0], device="cuda") + previous_read = torch.empty_like(live) + next_read = torch.empty_like(live) + source_stream = torch.cuda.Stream(device=live.device) + destination_stream = torch.cuda.Stream(device=live.device) + initial_stream = torch.cuda.current_stream(device=live.device) + original_overlap_stream = g_infer_context.overlap_stream + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.transfer = Transfer(live, staging) + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + manager.in_flight_layers = [0] + manager._commit_layer_metadata = lambda _layer: None + manager._finish_rebalance = lambda: None + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + + try: + g_infer_context.overlap_stream = source_stream + with torch.cuda.stream(source_stream): + source_stream.wait_stream(initial_stream) + torch.cuda._sleep(20_000_000) + previous_read.copy_(live, non_blocking=True) + with torch.cuda.stream(destination_stream): + manager._poll_in_flight() + with torch.cuda.stream(source_stream): + source_stream.wait_stream(destination_stream) + next_read.copy_(live, non_blocking=True) + source_stream.synchronize() + + assert previous_read.item() == 1.0 + assert next_read.item() == 2.0 + finally: + g_infer_context.overlap_stream = original_overlap_stream + + +def test_transfer_ring_reuses_a_buffer_only_after_commit_and_consumption(monkeypatch): + operations = [] + + class Event: + def __init__(self): + self.recorded = 0 + self.synchronized = 0 + + def record(self, stream): + self.recorded += 1 + + def synchronize(self): + self.synchronized += 1 + operations.append("consumed synchronize") + + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer.backend = "test" + transfer.device = torch.device("cuda", 0) + transfer.staging_depth = 2 + transfer.staging = [[], []] + transfer.live = [[], [], []] + transfer.num_experts_per_rank = 0 + transfer._release = [threading.Event(), threading.Event()] + for release in transfer._release: + release.set() + transfer._consumed_events = [Event(), Event()] + transfer._consumed_recorded = [False, False] + transfer._changed_dst_slots = [(), ()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = True + transfer.transfer_group = "transfer-group" + copied = [] + + def copy_batch(batch, _prepared_batch): + for layer, _plan, _buffer, _staging in batch: + copied.append(layer) + operations.append(("copy", layer)) + + transfer._copy_batch = copy_batch + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda device: None) + monkeypatch.setattr(transfer_module.torch.cuda, "Event", Event) + monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) + monkeypatch.setattr( + transfer_module.dist, + "barrier", + lambda **kwargs: operations.append(("barrier", kwargs["group"])), + ) + + plans = [(0, []), (1, []), (2, [])] + prepared_batches = transfer.prepare_transfer(plans) + monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) + transfer.start(plans, prepared_batches) + deadline = time.monotonic() + 2 + while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(0, 0), (1, 1)] + assert copied == [0, 1] + assert operations == [("copy", 0), ("copy", 1)] + + transfer.commit(0, 0) + while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: + time.sleep(0.001) + assert transfer.pending_layers() == [(1, 1), (2, 0)] + assert copied == [0, 1, 2] + assert transfer._consumed_events[0].synchronized == 1 + assert operations == [ + ("copy", 0), + ("copy", 1), + "consumed synchronize", + ("barrier", "transfer-group"), + ("copy", 2), + ] + transfer.commit(1, 1) + transfer.commit(2, 0) + transfer.finish() + + +def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once( + monkeypatch, +): + def make_transfer(finalize): + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer.backend = "test" + transfer.device = torch.device("cuda", 0) + transfer.global_rank = 0 + transfer.staging_depth = 1 + transfer.staging = [[]] + transfer._release = [threading.Event()] + transfer._release[0].set() + transfer._consumed_events = [object()] + transfer._consumed_recorded = [False] + transfer._changed_dst_slots = [()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = False + transfer._copy_batch = lambda _batch, _prepared_batch: None + transfer._finish_transfer_generation = finalize + return transfer + + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + finalized_before_publish = [] + success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) + success.start([(0, [])]) + success.finish() + assert finalized_before_publish == [0] + assert success.pending_layers() == [(0, 0)] + + failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) + failed.start([(0, [])]) + failed._thread.join() + assert list(failed._pending) == [] + with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: + failed.pending_layers() + assert isinstance(exc_info.value.__cause__, RuntimeError) + assert str(exc_info.value.__cause__) == "cache boom" + with pytest.raises(RuntimeError, match="EPLB migration worker failed"): + failed.finish() + assert failed._thread is None + + +def test_manager_rearms_after_rebalance_for_interval_one(): + recording_calls = [] + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.step_interval = 1 + manager.sampling_interval = 1 + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = 0 + manager.weights = [] + manager._eplb_states = [] + manager.target_placement = torch.zeros(1) + manager.in_flight_started_at = 0 + manager.global_rank = 1 + manager._set_recording = lambda enabled: recording_calls.append(enabled) + manager._finish_rebalance() + assert manager.in_flight is False + assert recording_calls == [True] + assert manager._steady_collection_end_step is None + assert manager._continuous_collection_start_step is None + + +def test_manager_sparse_insufficient_schedules_bounded_fresh_window(monkeypatch): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "insufficient", + "minimum_layer_samples": 1, + "minimum": 2, + } + manager.global_rank = 0 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.prefill_steps = 37 + manager.weights = [] + manager._eplb_states = [] + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + recordings, resets, logs = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) + + assert not manager._poll_evaluation() + assert not hasattr(manager, "_retained_local_samples") + assert manager._continuous_collection_start_step == 40 + assert manager._continuous_collection_end_step == 60 + assert recordings == [False] + assert resets == [True] + assert manager.sampling_interval == 20 + assert "insufficient samples" in logs[0][0] + assert "scheduled_fresh_window" in logs[0][0] + + manager.in_flight = False + manager.evaluation_in_flight = False + starts = [] + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) + manager.step() + manager.step() + assert manager.prefill_steps == 39 + assert starts == [] + manager.step() + assert manager.prefill_steps == 40 + assert starts == [] + for _ in range(19): + manager.step() + assert manager.prefill_steps == 59 + assert starts == [] + manager.step() + assert starts == [60] + + +def test_manager_full_window_insufficient_clears_and_backs_off(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "insufficient", + "minimum_layer_samples": 1, + "minimum": 2, + } + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.prefill_steps = 60 + manager.weights = [] + manager._eplb_states = [] + manager._continuous_collection_start_step = 40 + manager._continuous_collection_end_step = 60 + recordings, resets = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + + assert not manager._poll_evaluation() + assert not hasattr(manager, "_retained_local_samples") + assert manager._continuous_collection_start_step is None + assert manager._continuous_collection_end_step is None + assert manager.sampling_interval == 80 + assert recordings == [False] + assert resets == [True] + + +def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = False + manager.prefill_steps = 36 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager._steady_collection_end_step = manager.prefill_steps + 1 + recordings, resets, starts = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + + # The window is never truncated to the next boundary: it waits until 40, + # records a full 20 fresh steps, then evaluates at the boundary at 60. + manager._begin_continuous_collection() + assert manager._continuous_collection_start_step == 40 + assert manager._continuous_collection_end_step == 60 + assert recordings == [False] + assert resets == [True] + assert manager._steady_collection_end_step is None + + for _ in range(4): + manager.step() + assert manager.prefill_steps == 40 + assert recordings == [False, True] + assert starts == [] + for _ in range(19): + manager.step() + assert starts == [] + manager.step() # 60: the full window ends and triggers the evaluation. + assert starts == [True] + + +def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.prefill_steps = 80 + manager.step_interval = 20 + manager.sampling_interval = 80 + manager._steady_collection_end_step = None + recordings = [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: None + + manager._begin_continuous_collection() + assert manager._continuous_collection_start_step == 140 + assert manager._continuous_collection_end_step == 160 + assert recordings == [False] + + +def test_first_no_improvement_switches_to_sparse_sampling_window(monkeypatch): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager._continuous_collection_start_step = 0 + manager.weights = [] + manager._eplb_states = [] + recordings = [] + manager._set_recording = lambda enabled: recordings.append(enabled) + + assert not manager._poll_evaluation() + assert manager._continuous_collection_start_step is None + assert recordings == [False] + assert manager.sampling_interval == 80 + + +def test_continuous_collection_evaluates_only_after_one_full_base_window(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager._continuous_collection_start_step = 0 + manager._continuous_collection_end_step = 20 + manager.prefill_steps = 0 + manager.step_interval = 20 + manager.sampling_interval = 320 + manager._steady_collection_end_step = None + manager.evaluation_in_flight = False + started = [] + monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) + + for _ in range(19): + manager.step() + assert manager.prefill_steps == 19 + assert started == [] + + manager.step() + assert manager.prefill_steps == 20 + assert started == [True] + + +def test_no_improvement_exponentially_backs_off_sampling_interval_at_cap(): + class DoneThread: + def join(self): + pass + + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager._evaluation_lock = threading.Lock() + manager._set_recording = lambda _enabled: None + manager.global_rank = 1 + manager.step_interval = 20 + manager.sampling_interval = 20 + manager.weights = [] + manager._eplb_states = [] + + for expected_interval in (80, 320, 320): + manager.evaluation_in_flight = True + manager._evaluation_error = None + manager._evaluation_thread = DoneThread() + manager._evaluation_result = { + "kind": "no_improvement", + "model_imbalance_ratio": 1.2, + "candidate_model_imbalance_ratio": 1.1, + "candidate_rebalance_gain": 0.01, + "candidate_changed_layer_count": 2, + } + assert not manager._poll_evaluation() + assert manager.sampling_interval == expected_interval + + +def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.prefill_steps = 18 + manager.step_interval = 20 + manager.sampling_interval = 80 + manager._steady_collection_end_step = None + manager._continuous_collection_start_step = None + manager._continuous_collection_end_step = None + manager.evaluation_in_flight = False + recordings, resets, starts = [], [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._reset_recorded_samples = lambda: resets.append(True) + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + + manager.step() + manager.step() + assert manager.prefill_steps == 20 + assert recordings == [] + assert starts == [] + + for _ in range(55): + manager.step() + assert manager.prefill_steps == 75 + assert recordings == [] + assert starts == [] + + manager.step() + assert manager.prefill_steps == 76 + assert recordings == [True] + assert resets == [True] + assert manager._steady_collection_end_step is not None + + for _ in range(4): + manager.step() + assert manager.prefill_steps == 80 + assert starts == [True] + assert manager._steady_collection_end_step is None + + +def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.current_placement = torch.zeros((1, 1, 1), dtype=torch.int64) + manager.num_logical_experts = 1 + manager.world_size = 1 + manager.node_world_size = 1 + manager.step_interval = 20 + manager.sampling_interval = 320 + manager._continuous_collection_start_step = 0 + manager.global_rank = 1 + manager.transfer = type( + "Transfer", + (), + {"start": lambda self, plans, prepared_batches: setattr(self, "started", (plans, prepared_batches))}, + )() + manager._reset_recorded_samples = lambda: None + + prepared_batches = [object()] + manager._start_rebalance( + { + "placement": torch.zeros((1, 1, 1), dtype=torch.int64), + "improved": torch.tensor([True]), + "metadata": [None], + "layer_plans": [(0, object())], + "prepared_batches": prepared_batches, + "before": {"max": 1.0, "p95": 1.0}, + "after": {"max": 1.0, "p95": 1.0}, + "model_imbalance_ratio": 1.0, + "candidate_model_imbalance_ratio": 1.0, + "candidate_rebalance_gain": 0.1, + "candidate_changed_layer_count": 1, + } + ) + + assert manager.sampling_interval == 20 + assert manager.in_flight + assert manager._continuous_collection_start_step is None + assert len(manager.transfer.started[0]) == 1 + assert manager.transfer.started[1] is prepared_batches + + +def test_transfer_start_rejects_prepared_batches_with_wrong_batch_count(): + transfer = object.__new__(transfer_module._EPLBTransferBase) + transfer._thread = None + transfer.staging_depth = 2 + transfer.staging = [object(), object()] + + with pytest.raises(ValueError, match="prepared batch count"): + transfer.start([(0, []), (1, []), (2, [])], prepared_batches=[object()]) + + +def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.step_interval = 20 + manager.sampling_interval = manager.step_interval + manager._steady_collection_end_step = None + manager.weights = [] + manager._eplb_states = [] + manager.target_placement = torch.zeros(1) + manager.in_flight_started_at = 0 + manager.global_rank = 1 + recordings, starts = [], [] + manager._set_recording = lambda enabled: recordings.append(enabled) + manager._finish_rebalance() + assert recordings == [False] + + manager.in_flight = False + manager.prefill_steps = 38 + manager.evaluation_in_flight = False + monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) + manager.prefill_steps = 35 + manager.step() + assert recordings == [False, True] + for _ in range(4): + manager.step() + assert starts == [True] + + +def test_manager_inflight_step_does_not_poll(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = True + calls = [] + manager._poll_in_flight = lambda: calls.append("poll") + manager.step() + assert calls == [] + + +def test_manager_poll_advances_inflight_before_evaluation(): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = True + manager.evaluation_in_flight = True + calls = [] + manager._poll_in_flight = lambda: calls.append("inflight") + manager._poll_evaluation = lambda: calls.append("evaluation") + + manager.poll() + + assert calls == ["inflight"] + + +def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): + manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) + manager.in_flight = False + manager.evaluation_in_flight = True + manager._evaluation_lock = threading.Lock() + manager._evaluation_result = None + manager._evaluation_error = None + manager.control_group = object() + manager._control_ready_count = torch.empty(1, dtype=torch.int32) + calls = [] + manager._poll_evaluation = lambda: calls.append("evaluation") + + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) + manager.poll() + assert calls == [] + + manager._evaluation_result = {"kind": "no_improvement"} + monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) + manager.poll() + assert calls == ["evaluation"] + + +def test_nixl_descriptor_runs_merge_only_jointly_contiguous_source_and_destination_rows(): + steps = [ + TransferStep(0, 2, 1, 3), + TransferStep(0, 0, 1, 1), + TransferStep(0, 1, 1, 2), + TransferStep(0, 4, 1, 7), + ] + runs = transfer_module.NixlEPLBTransfer._contiguous_runs(steps) + assert [[(step.dst_slot, step.src_local_row) for step in run] for run in runs] == [ + [(0, 1), (1, 2), (2, 3)], + [(4, 7)], + ] + + +def test_nixl_prepare_batch_compiles_hot_path_without_tensor_views(monkeypatch): + class BatchMemcpy: + def __init__(self): + self.prepared = [] + self.enqueued = [] + + def prepare(self, descriptors): + descriptor = tuple(descriptors) + self.prepared.append(descriptor) + return descriptor + + def enqueue(self, descriptor, stream): + self.enqueued.append((descriptor, stream)) + + stream = SimpleNamespace(cuda_stream=123, synchronize=lambda: None) + batch_memcpy = BatchMemcpy() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._push_stream = stream + transfer._batch_memcpy = batch_memcpy + transfer.global_rank = 0 + transfer._same_node_ranks = {0, 1, 2} + transfer._live_row_layout = [ + [ + ("w13.weight", 1000, 32), + ("w13.weight_scale", 2000, 32), + ("w13.weight_zero_point", 3000, 32), + ] + ] + transfer._push_staging_row_layout = { + 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w13.weight_zero_point", 6000, 32)]], + 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w13.weight_zero_point", 9000, 32)]], + } + transfer._get_remote_read = lambda *_args: None + transfer._wait_xfers = lambda _xfers: None + + run_a = [TransferStep(1, 3, 0, 5), TransferStep(1, 4, 0, 6)] + run_b = [TransferStep(2, 1, 0, 2)] + local_inbound = TransferStep(0, 0, 1, 0) + remote_inbound = TransferStep(0, 1, 3, 2) + staging = object() + batch = [(0, run_a + run_b + [local_inbound, remote_inbound], 0, staging)] + prepared = transfer._prepare_batch(batch) + monkeypatch.setattr(transfer, "_prepare_batch", lambda _batch: pytest.fail("hot path must not prepare descriptors")) + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + transfer._copy_batch(batch, prepared) + + expected = ( + (1160, 4096, 64), + (2160, 5096, 64), + (3160, 6096, 64), + (1064, 7032, 32), + (2064, 8032, 32), + (3064, 9032, 32), + ) + assert batch_memcpy.prepared == [expected] + assert batch_memcpy.enqueued == [(expected, 123)] + assert prepared.remote_entries == {3: [(0, [remote_inbound], staging)]} + + +def test_nixl_prepare_transfer_batches_match_staging_depth(): + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.staging_depth = 2 + transfer.staging = ["staging-0", "staging-1"] + seen_batches = [] + + def prepare_batch(batch): + seen_batches.append(batch) + return f"prepared-{len(seen_batches)}" + + transfer._prepare_batch = prepare_batch + layer_plans = [(3, "plan-3"), (4, "plan-4"), (5, "plan-5")] + + prepared_batches = transfer.prepare_transfer(layer_plans) + + assert prepared_batches == [ + (seen_batches[0], "prepared-1"), + (seen_batches[1], "prepared-2"), + ] + assert [[(layer, buffer) for layer, _plan, buffer, _staging in batch] for batch in seen_batches] == [ + [(3, 0), (4, 1)], + [(5, 0)], + ] + + +def test_cuda_batch_memcpy_cuda13_abi_and_descriptor_layout(): + class Function: + def __init__(self, callback): + self.callback = callback + self.restype = None + self.argtypes = None + + def __call__(self, *args): + return self.callback(*args) + + class Library: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *args: self.calls.append(args) or 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + self.calls = [] + + import ctypes + + library = Library() + batch_memcpy = _CudaBatchMemcpy(library) + prepared = batch_memcpy.prepare(((101, 201, 64), (102, 202, 128))) + batch_memcpy.enqueue(prepared, 777) + + assert len(library.cudaMemcpyBatchAsync.argtypes) == 8 + assert library.calls[0][3] == 2 + assert library.calls[0][6] == 1 + assert library.calls[0][7].value == 777 + assert [pointer for pointer in library.calls[0][0]] == [201, 202] + assert [pointer for pointer in library.calls[0][1]] == [101, 102] + assert list(library.calls[0][2]) == [64, 128] + attrs = library.calls[0][4]._obj + assert attrs.srcAccessOrder == 1 + assert attrs.srcLocHint.type == attrs.srcLocHint.id == 0 + assert attrs.dstLocHint.type == attrs.dstLocHint.id == 0 + assert attrs.flags == 1 + + +def test_cuda_batch_memcpy_rejects_unsupported_runtime_and_invalid_descriptors(): + class Function: + def __init__(self, callback): + self.callback = callback + self.restype = None + self.argtypes = None + + def __call__(self, *args): + return self.callback(*args) + + class OldRuntimeLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 12080 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + import ctypes + + with pytest.raises(RuntimeError, match="13.0"): + _CudaBatchMemcpy(OldRuntimeLibrary()) + + class FutureRuntimeLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 14000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + with pytest.raises(RuntimeError, match="13.x"): + _CudaBatchMemcpy(FutureRuntimeLibrary()) + + class MissingBatchSymbolLibrary: + def __init__(self): + def get_version(pointer): + ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 + return 0 + + self.cudaRuntimeGetVersion = Function(get_version) + self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") + + with pytest.raises(RuntimeError, match="cudaMemcpyBatchAsync"): + _CudaBatchMemcpy(MissingBatchSymbolLibrary()) + with pytest.raises(ValueError, match="at least one"): + _CudaBatchMemcpy.prepare(()) + with pytest.raises(ValueError, match="positive"): + _CudaBatchMemcpy.prepare(((1, 2, 0),)) + + +def test_nixl_transfer_fails_fast_without_cuda13_batch_memcpy(monkeypatch): + failure = RuntimeError("missing cudaMemcpyBatchAsync") + + def unavailable(): + raise failure + + monkeypatch.setattr( + transfer_module._EPLBTransferBase, "__init__", lambda self, *_args: setattr(self, "device", "mock") + ) + monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) + monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", unavailable) + with pytest.raises(RuntimeError, match="missing cudaMemcpyBatchAsync") as exc_info: + transfer_module.NixlEPLBTransfer([object()], object(), 0, 1) + assert exc_info.value is failure + + +@pytest.mark.parametrize( + ("maps", "open_raises", "expected"), + [ + ( + "7f /cuda-12/libcudart.so.12.8\n" + "7f /first/libcudart.so.13 (deleted)\n" + "7f /second/libcudart.so.13\n" + "7f /cuda-14/libcudart.so.14\n" + "7f /cuda-130/libcudart.so.130\n", + False, + "/first/libcudart.so.13", + ), + ("7f /cuda-12/libcudart.so.12\n7f /cuda-130/libcudart.so.130\n", False, None), + ("", True, None), + ], +) +def test_cuda_batch_memcpy_finds_first_loaded_cuda13_runtime(monkeypatch, maps, open_raises, expected): + def open_maps(_path): + if open_raises: + raise OSError("maps unavailable") + return io.StringIO(maps) + + monkeypatch.setattr(builtins, "open", open_maps) + assert _CudaBatchMemcpy._find_loaded_cudart() == expected + + +@pytest.mark.parametrize(("layers", "expected_depth"), [(1, 1), (8, 8), (9, 8), (43, 8)]) +def test_nixl_transfer_bounds_staging_depth(monkeypatch, layers, expected_depth): + def base_init(self, *_args): + self.device = "mock-device" + + monkeypatch.setattr(transfer_module._EPLBTransferBase, "__init__", base_init) + monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) + monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", lambda: object()) + monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_ipc_metadata", lambda self: None) + monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_push_layouts", lambda self: None) + + transfer = transfer_module.NixlEPLBTransfer([object()] * layers, object(), 0, 1) + + assert transfer.staging_depth == expected_depth + + +def test_nixl_remote_read_cache_reuses_exact_batch_key_and_releases_on_shutdown(): + class Tensor: + nbytes = 8 + + def __init__(self, pointer): + self.pointer = pointer + + def data_ptr(self): + return self.pointer + + def get_device(self): + return 0 + + def __getitem__(self, _): + return self + + class Agent: + def __init__(self): + self.prepared = 0 + self.made = 0 + self.released_xfers = 0 + self.released_dlists = 0 + self.removed_agents = [] + + def get_xfer_descs(self, descriptors, _): + return descriptors + + def prep_xfer_dlist(self, *_args, **_kwargs): + self.prepared += 1 + return f"dlist-{self.prepared}" + + def make_prepped_xfer(self, *_args, **_kwargs): + self.made += 1 + return f"xfer-{self.made}" + + def query_xfer_backend(self, _): + return "UCX" + + def release_xfer_handle(self, _): + self.released_xfers += 1 + + def release_dlist_handle(self, _): + self.released_dlists += 1 + + def remove_remote_agent(self, remote_name): + self.removed_agents.append(remote_name) + + agent = Agent() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._nixl_agent = agent + transfer._xfer_cache = {} + transfer._used_xfer_cache_keys = set() + transfer._remote_agents = {1: "remote-1"} + transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} + transfer._registered_descs = None + transfer.live = [[("weight", Tensor(100))]] + staging = [("weight", Tensor(200))] + first_entries = [(0, [TransferStep(0, 0, 1, 0)], staging)] + + first = transfer._get_remote_read(1, first_entries) + assert transfer._get_remote_read(1, first_entries) is first + assert agent.made == 1 + + changed_entries = [(0, [TransferStep(0, 1, 1, 0)], staging)] + transfer._get_remote_read(1, changed_entries) + assert agent.made == 2 + + transfer.shutdown() + assert agent.released_xfers == 2 + assert agent.released_dlists == 4 + assert agent.removed_agents == ["remote-1"] + + +def test_nixl_remote_read_cache_is_bounded_to_the_current_transfer_generation( + monkeypatch, +): + class Tensor: + nbytes = 8 + + def __init__(self, pointer): + self.pointer = pointer + + def data_ptr(self): + return self.pointer + + def get_device(self): + return 0 + + def __getitem__(self, _): + return self + + class Agent: + def __init__(self): + self.made = 0 + self.released_xfers = 0 + self.released_dlists = 0 + + def get_xfer_descs(self, descriptors, _): + return descriptors + + def prep_xfer_dlist(self, *_args, **_kwargs): + return object() + + def make_prepped_xfer(self, *_args, **_kwargs): + self.made += 1 + return object() + + def query_xfer_backend(self, _): + return "UCX" + + def release_xfer_handle(self, _): + self.released_xfers += 1 + + def release_dlist_handle(self, _): + self.released_dlists += 1 + + def remove_remote_agent(self, _): + pass + + agent = Agent() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._nixl_agent = agent + transfer._xfer_cache = {} + transfer._used_xfer_cache_keys = set() + transfer._remote_agents = {1: "remote-1"} + transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} + transfer._registered_descs = None + transfer._ipc_staging = {} + transfer.live = [[("weight", Tensor(100))]] + transfer.device = torch.device("cuda", 0) + transfer.transfer_group = object() + transfer.global_rank = 0 + transfer.world_size = 2 + transfer.staging_depth = 1 + transfer.staging = [[]] + transfer.num_experts_per_rank = 0 + transfer._release = [threading.Event()] + transfer._release[0].set() + transfer._consumed_events = [object()] + transfer._consumed_recorded = [False] + transfer._changed_dst_slots = [()] + transfer._pending = deque() + transfer._pending_lock = threading.Lock() + transfer._error = None + transfer._thread = None + transfer._needs_staging_reuse_barrier = False + + staging = [("weight", Tensor(200))] + entries_a = [(0, [TransferStep(0, 0, 1, 0)], staging)] + entries_b = [(0, [TransferStep(0, 1, 1, 0)], staging)] + generation = [entries_a] + + def copy_batch(_batch, _prepared_batch): + transfer._get_remote_read(1, generation[0]) + + transfer._copy_batch = copy_batch + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + prepared_batches = [([(0, [], 0, transfer.staging[0])], None)] + transfer.start([(0, [])], prepared_batches) + transfer.finish() + assert agent.made == 1 + assert agent.released_xfers == agent.released_dlists == 0 + assert len(transfer._xfer_cache) == 1 + + # The real manager releases each staging buffer through commit(). This + # focused cache test has no commits, so model that hand-off before the + # next generation reuses buffer zero. + transfer._release[0].set() + transfer.start([(0, [])], prepared_batches) + transfer.finish() + assert agent.made == 1 + assert agent.released_xfers == agent.released_dlists == 0 + assert len(transfer._xfer_cache) == 1 + + generation[:] = [entries_b] + transfer._release[0].set() + transfer.start([(0, [])], prepared_batches) + transfer.finish() + assert agent.made == 2 + assert agent.released_xfers == 1 + assert agent.released_dlists == 2 + assert len(transfer._xfer_cache) == 1 + transfer.shutdown() + + +def test_nixl_ipc_metadata_exports_staging_per_local_target(monkeypatch): + from lightllm.server.router.model_infer.mode_backend.pd import p2p_fix + + class Tensor: + shape = (4, 2) + dtype = torch.float16 + device = torch.device("cuda", 0) + nbytes = 16 + + def __init__(self, label): + self.label = label + + def numel(self): + return 3 + + def __getitem__(self, _index): + return self + + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer.transfer_group = object() + transfer.global_rank = 0 + transfer.world_size = 3 + transfer.device = torch.device("cuda", 0) + transfer.live = [[("w13.weight", Tensor("local-w13")), ("w2.weight", Tensor("local-w2"))]] + transfer.staging_depth = 8 + transfer.staging = [[("w13.weight", Tensor(f"local-staging-{index}"))] for index in range(8)] + transfer._ipc_staging = {} + + reduce_calls, rebuild_calls, gathers = [], [], [] + + def reduce_tensor(tensor): + reduce_calls.append(tensor.label) + return None, (f"export-{tensor.label}",) + + def rebuild_tensor(export): + rebuild_calls.append(export) + return Tensor(export) + + source_one = [[("w13.weight", (4, 2), torch.float16, (f"rank1-staging-{index}",))] for index in range(8)] + + def all_gather(output, value, **_kwargs): + gathers.append(value) + if len(gathers) == 1: + output[:] = ["node-a", "node-a", "node-b"] + else: + output[:] = [value, {0: {"staging": source_one}}, {}] + + monkeypatch.setattr(p2p_fix, "reduce_tensor", reduce_tensor) + monkeypatch.setattr(p2p_fix, "p2p_fix_rebuild_cuda_tensor", rebuild_tensor) + monkeypatch.setattr(transfer_module.dist, "all_gather_object", all_gather) + monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) + + transfer._init_ipc_metadata() + + assert reduce_calls == [*(f"local-staging-{index}" for index in range(8))] + assert rebuild_calls == [*(f"rank1-staging-{index}" for index in range(8))] + assert transfer._same_node_ranks == {0, 1} + assert transfer._cross_node_ranks == {2} + assert transfer._needs_staging_reuse_barrier + assert [name for name, _ in transfer._ipc_staging[1][0]] == ["w13.weight"] + + +def test_nixl_copy_batch_source_pushes_local_rows_and_keeps_remote_ucx_reads( + monkeypatch, +): + class Stream: + def __init__(self): + self.synchronized = 0 + self.cuda_stream = 123 + + def synchronize(self): + self.synchronized += 1 + + stream = Stream() + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._push_stream = stream + transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) + remote_reads, waited_xfers, enqueued = [], [], [] + transfer._get_remote_read = lambda src_rank, entries: remote_reads.append((src_rank, entries)) or ( + None, + None, + "xfer", + ) + transfer._wait_xfers = waited_xfers.extend + + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + prepared_push = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") + transfer._copy_batch([], prepared_push) + assert enqueued == [("push", stream.cuda_stream)] + + remote_step = TransferStep(0, 1, 2, 0) + prepared_remote = transfer_module.NixlEPLBTransfer._PreparedBatch({2: [(0, [remote_step], [])]}, None) + transfer._copy_batch([], prepared_remote) + assert [rank for rank, _ in remote_reads] == [2] + assert waited_xfers == [(None, None, "xfer")] + + +def test_nixl_copy_batch_self_only_rank_pushes_and_synchronizes(monkeypatch): + class Stream: + def __init__(self): + self.synchronized = 0 + self.cuda_stream = 456 + + def synchronize(self): + self.synchronized += 1 + + transfer = object.__new__(transfer_module.NixlEPLBTransfer) + transfer._push_stream = Stream() + enqueued = [] + transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) + transfer._wait_xfers = lambda _xfers: None + monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) + + prepared = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") + transfer._copy_batch([], prepared) + + assert enqueued == [("push", transfer._push_stream.cuda_stream)] + assert transfer._push_stream.synchronized == 1 + + +def test_manager_constructs_nixl_transfer(monkeypatch): + weight = type( + "Weight", + (), + { + "n_routed_experts": 4, + "expert_parallel_state": _test_parallel_state( + eplb=True, + num_logical_experts=4, + world_size=2, + num_redundant_experts_per_rank=2, + route_counter=torch.zeros((2, 4), dtype=torch.int64), + ), + }, + )() + transfer = object() + groups = [object(), object(), object()] + new_group_calls = [] + monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) + monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) + monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) + monkeypatch.setattr(manager_module, "get_node_world_size", lambda: 2) + monkeypatch.setattr(manager_module, "get_prefill_eplb_step_interval", lambda: 20) + monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) + + def new_group(*args, **kwargs): + new_group_calls.append((args, kwargs)) + return groups[len(new_group_calls) - 1] + + monkeypatch.setattr(manager_module.dist, "new_group", new_group) + transfer_calls = [] + monkeypatch.setattr( + manager_module, + "NixlEPLBTransfer", + lambda weights, group, rank, world_size: ( + transfer_calls.append((weights, group, rank, world_size)) or transfer + ), + ) + logs = [] + monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) + manager = manager_module.EPLBManager(type("Model", (), {})()) + assert manager.transfer is transfer + assert ( + manager.evaluation_group, + manager.control_group, + manager.transfer_group, + ) == tuple(groups) + assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 + assert transfer_calls == [([weight], groups[2], 0, 2)] + assert manager.rebalance_gain_threshold == 0.07 + assert "rebalance_gain_threshold=0.0700" in logs[0] + assert manager._continuous_collection_start_step is None + assert manager._continuous_collection_end_step == manager.step_interval + assert weight.expert_parallel_state.eplb.recording + assert manager._eplb_states[0] is weight.expert_parallel_state.eplb + assert not hasattr(weight.expert_parallel_state.eplb, "record_load") + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("tokens", [1, 32]) +@pytest.mark.parametrize("scoring_func", ["sigmoid", "softmax"]) +@pytest.mark.parametrize("renormalize", [False, True]) +def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens, scoring_func, renormalize): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import ( + triton_grouped_topk, + triton_grouped_topk_eplb, + ) + + torch.manual_seed(1234) + topk = 8 + experts = 256 + num_expert_group = 8 + gating_output = torch.randn((tokens, experts), dtype=torch.bfloat16, device="cuda") + correction_bias = torch.randn((experts,), dtype=torch.float32, device="cuda") + hidden_states = torch.empty((tokens, 1), dtype=torch.bfloat16, device="cuda") + logical_to_physical = torch.stack( + ( + torch.arange(experts, dtype=torch.int32, device="cuda"), + torch.arange(experts, dtype=torch.int32, device="cuda") + experts, + ), + dim=1, + ) + logical_replica_count = torch.where( + torch.arange(experts, device="cuda") % 3 == 0, + torch.full((experts,), 2, dtype=torch.int32, device="cuda"), + torch.ones((experts,), dtype=torch.int32, device="cuda"), + ) + expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + fused_counter = torch.zeros_like(expected_counter) + + expected_weights, logical_ids = triton_grouped_topk( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group, + 4, + scoring_func, + 2, + ) + if tokens == 1: + replica_indices = torch.zeros_like(logical_ids) + else: + token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + replica_indices = ( + (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) + & 0xFFFFFFFF + ) % logical_replica_count[logical_ids.to(torch.long)].to(torch.int64) + expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)].to(torch.long) + if record_load: + expected_counter[1].scatter_add_( + 0, + logical_ids.reshape(-1).to(torch.long), + torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), + ) + fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( + hidden_states, + gating_output, + correction_bias, + topk, + renormalize, + num_expert_group, + 4, + scoring_func, + logical_to_physical, + logical_replica_count, + fused_counter, + sample_index=1, + record_load=record_load, + use_grouped_topk=True, + group_score_used_topk_num=2, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(fused_weights, expected_weights, rtol=1e-5, atol=1e-6) + assert torch.equal(fused_ids, expected_ids) + assert fused_logical_ids is None + assert torch.equal(fused_counter, expected_counter) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("tokens", [1, 32]) +def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb + + torch.manual_seed(1234) + topk = 4 + experts = 64 + gating_output = torch.randn((tokens, experts), dtype=torch.float32, device="cuda") + logical_to_physical = torch.stack( + ( + torch.arange(experts, dtype=torch.int32, device="cuda"), + torch.arange(experts, dtype=torch.int32, device="cuda") + experts, + ), + dim=1, + ) + logical_replica_count = torch.where( + torch.arange(experts, device="cuda") % 3 == 0, + torch.full((experts,), 2, dtype=torch.int32, device="cuda"), + torch.ones((experts,), dtype=torch.int32, device="cuda"), + ) + expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") + fused_counter = torch.zeros_like(expected_counter) + expected_weights, expected_logical_ids = torch.softmax(gating_output, dim=-1).topk(topk, dim=-1) + if tokens == 1: + replica_indices = torch.zeros_like(expected_logical_ids) + else: + token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) + replica_indices = ( + ( + ((token_indices * 2654435769) & 0xFFFFFFFF) + + ((expected_logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF) + ) + & 0xFFFFFFFF + ) % logical_replica_count[expected_logical_ids].to(torch.int64) + expected_ids = logical_to_physical[expected_logical_ids, replica_indices.to(torch.long)].to(torch.long) + if record_load: + expected_counter[1].scatter_add_( + 0, + expected_logical_ids.reshape(-1), + torch.ones(expected_logical_ids.numel(), dtype=torch.int64, device="cuda"), + ) + expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) + + weights, physical_ids, logical_ids = triton_grouped_topk_eplb( + hidden_states=torch.empty((tokens, 1), dtype=torch.float32, device="cuda"), + gating_output=gating_output, + correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), + topk=topk, + renormalize=True, + num_expert_group=8, + topk_group=4, + scoring_func="sigmoid", + logical_to_physical_map=logical_to_physical, + logical_replica_count=logical_replica_count, + expert_counter=fused_counter, + sample_index=1, + record_load=record_load, + use_grouped_topk=False, + return_logical_ids=True, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(weights, expected_weights, rtol=1e-5, atol=1e-6) + assert torch.equal(physical_ids, expected_ids) + assert torch.equal(logical_ids, expected_logical_ids) + assert torch.equal(fused_counter, expected_counter) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") +def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): + from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb + + experts = 64 + counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") + weights, physical_ids, logical_ids = triton_grouped_topk_eplb( + hidden_states=torch.empty((0, 1), device="cuda"), + gating_output=torch.empty((0, experts), device="cuda"), + correction_bias=None, + topk=4, + renormalize=False, + num_expert_group=8, + topk_group=4, + scoring_func="softmax", + logical_to_physical_map=torch.zeros((experts, 1), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones((experts,), dtype=torch.int32, device="cuda"), + expert_counter=counter, + sample_index=0, + record_load=True, + use_grouped_topk=False, + return_logical_ids=True, + ) + + assert weights.shape == physical_ids.shape == logical_ids.shape == (0, 4) + assert physical_ids.dtype is logical_ids.dtype is torch.long + assert torch.equal(counter, torch.zeros_like(counter)) diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py new file mode 100644 index 0000000000..46c218ad45 --- /dev/null +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -0,0 +1,444 @@ +"""NIXL EPLB correctness tests and a two-GPU 512 MiB micro-performance test.""" +import os +import random +import socket +import statistics +import time + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp + +from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( + NixlEPLBTransfer, + align_target_placement, + build_transfer_plan, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( + build_initial_redundant_expert_ids, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + EPLBState, + ExpertParallelState, +) + +pytest.importorskip("nixl", reason="NIXL package is required") + + +class _Pack: + def __init__(self, weight, weight_scale): + self.weight = weight + self.weight_scale = weight_scale + self.weight_zero_point = None + + +def _free_port(): + sock = socket.socket() + sock.bind(("127.0.0.1", 0)) + port = sock.getsockname()[1] + sock.close() + return port + + +class _FakeWeight: + def __init__(self, rank, layer_index, row_elements): + self.n_routed_experts = 32 + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=32, + world_size=2, + eplb=EPLBState( + num_redundant_experts_per_rank=16, + initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids(32, 2, 16), + logical_to_physical_map=torch.zeros((32, 2), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones(32, dtype=torch.int32, device="cuda"), + route_counter=torch.zeros((1, 32), dtype=torch.int64, device="cuda"), + ), + ) + base = rank * 100 + layer_index * 100 + self.w13 = self._pack(base, row_elements) + self.w2 = self._pack(base + 10, row_elements) + + @staticmethod + def _pack(base, row_elements): + weight = torch.empty((32, row_elements), dtype=torch.float16, device="cuda") + for row in range(weight.shape[0]): + weight[row].fill_(base + row) + scale = torch.empty((32, 1), dtype=torch.float32, device="cuda") + for row in range(scale.shape[0]): + scale[row].fill_(base + row + 0.5) + return _Pack(weight, scale) + + +def _wait_for_ready_prefix(transfer, control_group): + deadline = time.monotonic() + 30 + while True: + pending = transfer.pending_layers() + ready_count = torch.tensor([len(pending)], dtype=torch.int32) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) + if int(ready_count.item()) > 0: + return pending[: int(ready_count.item())] + if time.monotonic() >= deadline: + raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") + time.sleep(0.001) + + +def _run_layers( + transfer, + control_group, + layer_plans, + callback=lambda layer_index: None, + before_commit_callback=lambda layer_index: None, +): + transfer.start(layer_plans) + committed = 0 + while committed < len(layer_plans): + pending = _wait_for_ready_prefix(transfer, control_group) + assert len(pending) <= len(layer_plans) - committed + for layer_index, buffer_index in pending: + assert layer_index == layer_plans[committed][0] + before_commit_callback(layer_index) + transfer.commit( + layer_index, + buffer_index, + lambda layer_index=layer_index: callback(layer_index), + ) + committed += 1 + transfer.finish() + + +def _assert_correctness(weights, rank, source_rows_by_dst_slot): + if rank == 0: + for layer_index in (0, len(weights) - 1): + base = 100 + layer_index * 100 + for dst_slot, src_row in enumerate(source_rows_by_dst_slot): + dst_row = 16 + dst_slot + assert torch.all(weights[layer_index].w13.weight[dst_row] == base + src_row) + assert torch.all(weights[layer_index].w13.weight_scale[dst_row] == base + src_row + 0.5) + assert torch.all(weights[layer_index].w2.weight[dst_row] == base + src_row + 10) + assert torch.all(weights[layer_index].w2.weight_scale[dst_row] == base + src_row + 10.5) + + +def _benchmark(transfer, control_group, layer_plans, payload): + for _ in range(3): + _run_layers(transfer, control_group, layer_plans) + dist.barrier(group=control_group) + started = time.perf_counter() + for _ in range(8): + _run_layers(transfer, control_group, layer_plans) + torch.cuda.synchronize() + return payload * 8 / (time.perf_counter() - started) / 1e9 + + +def _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload): + batch = [ + (layer_index, plan, buffer_index, transfer.staging[buffer_index]) + for buffer_index, (layer_index, plan) in enumerate(layer_plans) + ] + # Measure the precompiled hot path only; planning and descriptor construction are off the timer. + prepared_batch = transfer._prepare_batch(batch) + for _ in range(3): + transfer._copy_batch(batch, prepared_batch) + dist.barrier(group=control_group) + samples = [] + for _ in range(20): + started = time.perf_counter() + transfer._copy_batch(batch, prepared_batch) + samples.append(payload / (time.perf_counter() - started) / 1e9) + torch.cuda.synchronize() + median = statistics.median(samples) + print( + f"NIXL _copy_batch payload={payload / 2**20:.1f} MiB; " + f"min={min(samples):.2f} GB/s median={median:.2f} " + f"mean={statistics.mean(samples):.2f} max={max(samples):.2f}", + flush=True, + ) + return median + + +def _eplb_worker(rank, port, queue): + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + torch.cuda.set_device(rank) + dist.init_process_group("gloo", rank=rank, world_size=2) + control_group = dist.new_group([0, 1], backend="gloo") + transfer_group = dist.new_group([0, 1], backend="gloo") + # Eight layers × 16 changed experts × two 2 MiB rows = 512 MiB useful remote weight payload. + row_elements = int(os.getenv("LIGHTLLM_EPLB_TEST_ROW_ELEMENTS", str(1024 * 1024))) + current = torch.tensor([list(range(16)), list(range(16))]) + source_rows_by_dst_slot = list(range(16)) + random.Random(20260731).shuffle(source_rows_by_dst_slot) + # Rank 0 receives the reverse logical-expert range in a deterministic random slot order. + # Consequently every descriptor has a distinct source and destination row. + target = torch.tensor([[16 + source_row for source_row in source_rows_by_dst_slot], list(range(16))]) + plan = build_transfer_plan(current, target, 32, 2, 2) + assert [step.src_local_row for step in plan if step.dst_rank == 0] == source_rows_by_dst_slot + benchmark_layer_count = 8 + layer_count = benchmark_layer_count + 1 + row_payload = ( + 2 * row_elements * torch.empty((), dtype=torch.float16).element_size() + + 2 * torch.empty((), dtype=torch.float32).element_size() + ) + payload = benchmark_layer_count * 16 * row_payload + weights = [_FakeWeight(rank, layer_index, row_elements) for layer_index in range(layer_count)] + transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=2) + assert transfer.staging_depth == 8 + assert transfer._eplb_states[0] is weights[0].expert_parallel_state.eplb + assert all(tensor.is_cuda for staging in transfer.staging for _, tensor in staging) + + wrap_layer_plans = [(layer_index, plan) for layer_index in range(layer_count)] + + def delay_rank_zero_first_commit(layer_index): + if rank == 0 and layer_index == 0: + time.sleep(0.1) + + _run_layers( + transfer, + control_group, + wrap_layer_plans, + before_commit_callback=delay_rank_zero_first_commit, + ) + torch.cuda.synchronize() + _assert_correctness(weights, rank, source_rows_by_dst_slot) + layer_plans = wrap_layer_plans[:benchmark_layer_count] + nixl_copy_batch = _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload) + nixl_bandwidth = _benchmark(transfer, control_group, layer_plans, payload) + transfer.shutdown() + + gathered = [None, None] + dist.all_gather_object(gathered, (nixl_bandwidth, nixl_copy_batch), group=control_group) + if rank == 0: + queue.put((payload, *gathered[0])) + dist.barrier(group=control_group) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires two CUDA GPUs", +) +def test_eplb_transfer_two_gpu_correctness_and_microperf(): + queue = mp.get_context("spawn").SimpleQueue() + mp.spawn(_eplb_worker, args=(_free_port(), queue), nprocs=2, join=True) + payload, nixl_gbps, nixl_copy_batch_gbps = queue.get() + print( + f"EPLB remote payload/round: {payload / 2**20:.1f} MiB; " + f"NIXL={nixl_gbps:.2f} GB/s NIXL _copy_batch={nixl_copy_batch_gbps:.2f} GB/s" + ) + assert nixl_gbps > 0 + + +def _depth_value(expert, layer_index, offset): + return expert // 32 * 100 + layer_index * 5 + expert % 32 + offset + + +class _DepthWeight: + def __init__(self, rank, layer_index, initial_placement): + self.n_routed_experts = 256 + self.expert_parallel_state = ExpertParallelState( + num_logical_experts=256, + world_size=8, + eplb=EPLBState( + num_redundant_experts_per_rank=4, + initial_redundant_expert_ids_by_rank=initial_placement.clone(), + logical_to_physical_map=torch.zeros((256, 8), dtype=torch.int32, device="cuda"), + logical_replica_count=torch.ones(256, dtype=torch.int32, device="cuda"), + route_counter=torch.zeros((1, 256), dtype=torch.int64, device="cuda"), + ), + ) + logical_ids = list(range(rank * 32, (rank + 1) * 32)) + initial_placement[rank].tolist() + self.w13 = self._pack(logical_ids, layer_index, 0) + self.w2 = self._pack(logical_ids, layer_index, 2) + + @staticmethod + def _pack(logical_ids, layer_index, offset): + weight = torch.empty((36, 64), dtype=torch.float16, device="cuda") + scale = torch.empty((36, 1), dtype=torch.float32, device="cuda") + for row, expert in enumerate(logical_ids): + value = _depth_value(expert, layer_index, offset) + weight[row].fill_(value) + scale[row].fill_(value + 0.25) + return _Pack(weight, scale) + + +def _depth_target(layer_index): + return torch.tensor([[((dst + layer_index + slot + 1) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]) + + +def _wait_all_pending(transfer, group, expected_count): + deadline = time.monotonic() + 30 + while True: + pending = transfer.pending_layers() + ready_count = torch.tensor([len(pending)], dtype=torch.int32) + dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=group) + if int(ready_count.item()) == expected_count: + return pending + if time.monotonic() > deadline: + raise TimeoutError(f"expected {expected_count} pending layers, got {pending}") + time.sleep(0.001) + + +def _clone_depth_live(weights): + return [ + [tensor.detach().clone() for _, tensor in transfer_tensors] + for transfer_tensors in [ + [ + ("w13.weight", weight.w13.weight), + ("w13.scale", weight.w13.weight_scale), + ("w2.weight", weight.w2.weight), + ("w2.scale", weight.w2.weight_scale), + ] + for weight in weights + ] + ] + + +def _assert_depth_snapshot(weights, snapshot, layer_indices=None, primary_only=False): + if layer_indices is None: + layer_indices = range(len(weights)) + for layer_index in layer_indices: + live_tensors = ( + weights[layer_index].w13.weight, + weights[layer_index].w13.weight_scale, + weights[layer_index].w2.weight, + weights[layer_index].w2.weight_scale, + ) + for live, expected in zip(live_tensors, snapshot[layer_index]): + if primary_only: + live = live[:32] + expected = expected[:32] + torch.testing.assert_close(live, expected) + + +def _assert_depth_staging(rank, layer_plans, target_placements, pending, transfer): + for buffer_index, ((layer_index, plan), pending_item) in enumerate(zip(layer_plans, pending)): + assert pending_item == (layer_index, buffer_index) + expected = {step.dst_slot: step for step in plan if step.dst_rank == rank} + for dst_slot, step in expected.items(): + expert = int(target_placements[layer_index][rank, dst_slot]) + base = _depth_value(expert, layer_index, 0) + staging = transfer.staging[buffer_index] + assert torch.all(staging[0][1][dst_slot] == base) + assert torch.all(staging[1][1][dst_slot] == base + 0.25) + assert torch.all(staging[2][1][dst_slot] == base + 2) + assert torch.all(staging[3][1][dst_slot] == base + 2.25) + + +def _assert_depth_live(weights, rank, layer_plans, target_placements): + for layer_index, plan in layer_plans: + for step in plan: + if step.dst_rank != rank: + continue + expert = int(target_placements[layer_index][rank, step.dst_slot]) + base = _depth_value(expert, layer_index, 0) + assert torch.all(weights[layer_index].w13.weight[32 + step.dst_slot] == base) + assert torch.all(weights[layer_index].w13.weight_scale[32 + step.dst_slot] == base + 0.25) + assert torch.all(weights[layer_index].w2.weight[32 + step.dst_slot] == base + 2) + assert torch.all(weights[layer_index].w2.weight_scale[32 + step.dst_slot] == base + 2.25) + + +def _assert_peer_coverage(layer_plans, require_redundant_source=False): + steps = [step for _, plan in layer_plans for step in plan] + assert {step.dst_rank for step in steps} == set(range(8)) + assert {step.src_rank for step in steps} == set(range(8)) + for source_rank in range(8): + assert len({step.dst_rank for step in steps if step.src_rank == source_rank}) > 1 + if require_redundant_source: + assert any(step.src_local_row >= 32 for step in steps) + + +def _depth_worker(rank, port): + os.environ["MASTER_ADDR"] = "127.0.0.1" + os.environ["MASTER_PORT"] = str(port) + torch.cuda.set_device(rank) + dist.init_process_group("gloo", rank=rank, world_size=8) + control_group = dist.new_group(list(range(8)), backend="gloo") + transfer_group = dist.new_group(list(range(8)), backend="gloo") + initial_placement = build_initial_redundant_expert_ids(256, 8, 4) + weights = [_DepthWeight(rank, layer_index, initial_placement) for layer_index in range(9)] + transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=8) + assert transfer.staging_depth == 8 + staging_bytes = sum(tensor.nbytes for staging in transfer.staging for _, tensor in staging) + one_layer_staging_bytes = sum(tensor[:4].nbytes for _, tensor in transfer.live[0]) + assert staging_bytes == 8 * one_layer_staging_bytes + current = initial_placement + first_order = [8, 0, 5, 1, 7, 3, 6, 2, 4] + first_targets = { + layer_index: align_target_placement(current, _depth_target(layer_index)) for layer_index in range(9) + } + first_plans = [ + (layer_index, build_transfer_plan(current, first_targets[layer_index], 256, 8, 8)) + for layer_index in first_order + ] + _assert_peer_coverage(first_plans) + first_snapshot = _clone_depth_live(weights) + transfer.start(first_plans, transfer.prepare_transfer(first_plans)) + committed = 0 + pending = _wait_all_pending(transfer, control_group, 8) + _assert_depth_staging(rank, first_plans[:8], first_targets, pending, transfer) + _assert_depth_snapshot(weights, first_snapshot) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) + transfer.commit(layer_index, buffer_index) + committed += 1 + pending = _wait_all_pending(transfer, control_group, 1) + _assert_depth_staging(rank, first_plans[8:], first_targets, pending, transfer) + _assert_depth_snapshot(weights, first_snapshot, [first_plans[8][0]]) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) + transfer.commit(layer_index, buffer_index) + committed += 1 + transfer.finish() + torch.cuda.synchronize() + dist.barrier(group=control_group) + _assert_depth_live(weights, rank, first_plans, first_targets) + + second_order = [7, 2, 4] + second_targets = { + layer_index: align_target_placement( + first_targets[layer_index], + torch.tensor([[((dst + layer_index + slot + 3) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]), + ) + for layer_index in second_order + } + second_plans = [ + ( + layer_index, + build_transfer_plan(first_targets[layer_index], second_targets[layer_index], 256, 8, 8), + ) + for layer_index in second_order + ] + _assert_peer_coverage(second_plans, require_redundant_source=True) + second_snapshot = _clone_depth_live(weights) + transfer.start(second_plans, transfer.prepare_transfer(second_plans)) + pending = _wait_all_pending(transfer, control_group, len(second_plans)) + _assert_depth_staging(rank, second_plans, second_targets, pending, transfer) + _assert_depth_snapshot(weights, second_snapshot) + dist.barrier(group=control_group) + if rank == 0: + time.sleep(0.1) + for layer_index, buffer_index in pending: + transfer.commit(layer_index, buffer_index) + transfer.finish() + torch.cuda.synchronize() + dist.barrier(group=control_group) + _assert_depth_live(weights, rank, second_plans, second_targets) + _assert_depth_snapshot(weights, second_snapshot, set(range(9)) - set(second_order)) + _assert_depth_snapshot(weights, second_snapshot, primary_only=True) + transfer.shutdown() + dist.barrier(group=control_group) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 8, + reason="requires eight CUDA GPUs", +) +def test_eplb_transfer_eight_gpu_bounded_staging_reuse(): + mp.spawn(_depth_worker, args=(_free_port(),), nprocs=8, join=True) diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py new file mode 100644 index 0000000000..95d41d521b --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_memory_profile.py @@ -0,0 +1,65 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager +from lightllm.utils import profile_max_tokens + + +@pytest.mark.parametrize("exclusion,expected", [(0, 1000), (200, 800), (None, 1000)]) +def test_mtp_profile_exclusion_adjustment(monkeypatch, exclusion, expected): + seen = [] + values = iter((100, 1100)) + monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) + monkeypatch.setattr(profile_max_tokens, "get_mtp_weight_layer_num", lambda: 1) + monkeypatch.setattr( + profile_max_tokens, "get_mtp_adjusted_mem_fraction", lambda **kw: seen.append(kw["target_weight_bytes"]) or 0.5 + ) + attrs = dict( + max_total_token_num=None, + is_mtp_draft_model=False, + args=SimpleNamespace(mtp_mode="x"), + config={"n_layer": 1}, + mem_fraction=0.8, + ) + if exclusion is not None: + attrs["get_mtp_profile_weight_exclusion"] = lambda: exclusion + model = SimpleNamespace(**attrs) + with profile_max_tokens.profile_mtp_weight_memory(model): + pass + assert seen == [expected] + + +@pytest.mark.parametrize("exclusion", [-1, 1001]) +def test_mtp_profile_exclusion_validation(monkeypatch, exclusion): + values = iter((100, 1100)) + monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) + model = SimpleNamespace( + max_total_token_num=None, + is_mtp_draft_model=False, + args=SimpleNamespace(mtp_mode="x"), + config={"n_layer": 1}, + mem_fraction=0.8, + get_mtp_profile_weight_exclusion=lambda: exclusion, + ) + with pytest.raises(ValueError, match="invalid MTP profile exclusion"): + with profile_max_tokens.profile_mtp_weight_memory(model): + pass + + +@pytest.mark.parametrize("reservations,expected", [({}, 252), ({"x": 20}, 247)]) +def test_memory_manager_profile_reservation_once(monkeypatch, reservations, expected): + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.torch.cuda.empty_cache", lambda: None) + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.dist.get_world_size", lambda: 1) + monkeypatch.setattr( + "lightllm.common.kv_cache_mem_manager.mem_manager.get_available_gpu_memory", lambda w: 1024 / 1024 ** 3 + ) + monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.get_total_gpu_memory", lambda: 0) + m = MemoryManager.__new__(MemoryManager) + m.size = None + m.memory_reservations = reservations + m.get_cell_size = lambda: 4 + m.get_fixed_memory_size = lambda: 16 + m.profile_size(1) + assert m.size == expected diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py new file mode 100644 index 0000000000..d566a94129 --- /dev/null +++ b/unit_tests/server/test_api_start_eplb.py @@ -0,0 +1,64 @@ +import pytest + +from lightllm.server import api_start +from lightllm.server.core.objs.start_args_type import StartArgs + + +def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeypatch): + args = StartArgs( + enable_ep_moe=True, + enable_prefill_eplb=True, + enable_prefill_cudagraph=True, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + monkeypatch.setattr( + api_start.process_manager, + "start_submodule_processes", + lambda *args, **kwargs: pytest.fail("subprocess startup must not be reached"), + ) + + with pytest.raises(AssertionError, match="--enable_prefill_eplb does not support --enable_prefill_cudagraph"): + api_start._launch_subprocesses(args) + + +def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monkeypatch): + args = StartArgs( + model_dir="test-model", + enable_ep_moe=True, + enable_prefill_eplb=True, + mtp_mode="vanilla_no_att", + mtp_step=1, + eos_id=0, + data_type="float16", + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) + monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) + monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) + monkeypatch.setattr(api_start, "get_model_type", lambda model_dir: "llama") + monkeypatch.setattr(api_start, "auto_set_response_parsers", lambda args: None) + monkeypatch.setattr(api_start, "auto_configure_allreduce_flags_from_args", lambda args: None) + monkeypatch.setattr(api_start, "validate_ports", lambda ports: None) + monkeypatch.setattr(api_start, "set_env_start_args", lambda args: None) + monkeypatch.setattr(api_start, "get_shm_port_args", lambda create=False: None) + monkeypatch.setattr(api_start, "send_and_receive_node_ip", lambda args: None) + monkeypatch.setattr(api_start, "is_sm100_gpu", lambda: False) + monkeypatch.setattr( + api_start.process_manager, + "start_submodule_processes", + lambda *args, **kwargs: (object(), None), + ) + monkeypatch.setattr(api_start.process_manager, "register_process_tree", lambda process: None) + + api_start._launch_subprocesses(args) From 33cd0c9d13f56cee332b533058c19685c923ecc9 Mon Sep 17 00:00:00 2001 From: niushengxiao Date: Sat, 29 Aug 2026 18:48:39 +0800 Subject: [PATCH 169/214] feat: support EPLB for DeepSeek v4 --- .../meta_weights/fused_moe/eplb_placement.py | 82 ++----- .../fused_moe/fused_moe_weight.py | 3 +- .../fused_moe/impl/deepgemm_impl.py | 6 - .../triton_kernel/fused_moe/grouped_topk.py | 1 - lightllm/common/eplb_utils.py | 2 +- .../deepseek4_mem_manager.py | 12 +- .../layer_infer/transformer_layer_infer.py | 57 ++++- lightllm/models/deepseek_v4/model.py | 85 +++++++ .../triton_kernel/csrc/moe_topk_eplb.cu | 219 ++++++++++++++++++ .../deepseek_v4/triton_kernel/moe_topk.py | 81 +++++++ .../model_infer/mode_backend/eplb_manager.py | 12 +- .../model_infer/mode_backend/eplb_transfer.py | 22 +- lightllm/utils/envs_utils.py | 1 + unit_tests/common/fused_moe/test_eplb.py | 107 +++++---- .../fused_moe/test_eplb_transfer_gpu.py | 2 +- .../models/deepseek_v4/test_memory_profile.py | 49 ++++ .../models/deepseek_v4/test_moe_topk.py | 91 ++++++++ .../deepseek_v4/test_vision_integration.py | 22 +- 18 files changed, 692 insertions(+), 162 deletions(-) create mode 100644 lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu create mode 100644 lightllm/models/deepseek_v4/triton_kernel/moe_topk.py create mode 100644 unit_tests/models/deepseek_v4/test_moe_topk.py diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py index d550355b9f..1ea8eb1643 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py @@ -97,11 +97,8 @@ def select_improving_placements( assert current_placement.shape == candidate_placement.shape current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) - current_critical = current_rank_load.max(dim=-1).values - candidate_critical = candidate_rank_load.max(dim=-1).values - if current_rank_load.ndim == 3: - current_critical = current_critical.sum(dim=0) - candidate_critical = candidate_critical.sum(dim=0) + current_critical = current_rank_load.max(dim=-1).values.sum(dim=0) + candidate_critical = candidate_rank_load.max(dim=-1).values.sum(dim=0) # Each changed layer must reduce its own critical load. All selected # changes must then collectively meet the configured model-level # critical-load reduction threshold, avoiding low-gain migrations. @@ -144,31 +141,28 @@ def plan_redundant_experts( current_placement: torch.Tensor | None = None, stickiness: float = 0.0, ) -> torch.Tensor: - """Plan replicas using source-node-local copies, with global fallback. + """Plan replicas from [samples, layers, source_nodes, experts] loads. - With ``current_placement`` and a positive ``stickiness``, a candidate that + With ``current_placement`` and positive ``stickiness``, a candidate that keeps an expert on its current rank receives a bonus of ``stickiness * mean per-layer expert load``. This preserves rank membership, not a particular redundant physical slot; target slots are canonicalized against the current live rows before transfer and metadata publication. A rank membership only changes when the move improves the critical-load objective by more than that margin. - Without them the planning is bit-identical to the legacy behavior. + With zero stickiness, placement is determined solely by the load objective. """ - assert expert_load.ndim in (2, 3, 4) if expert_alignment is not None: assert expert_alignment > 0 - use_legacy_topology_preference = expert_load.ndim < 4 - legacy_node_world_size = node_world_size if use_legacy_topology_preference else None - source_load, _squeeze_sample, node_world_size = _as_source_node_load(expert_load, num_ranks, node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + node_world_size = _resolve_node_world_size(expert_load, num_ranks, node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape assert num_logical_experts % num_ranks == 0 assert num_redundant_experts_per_rank > 0 num_experts_per_rank = num_logical_experts // num_ranks num_redundant = num_ranks * num_redundant_experts_per_rank assert num_redundant <= num_logical_experts * (num_ranks - 1) - load = source_load.to(dtype=torch.float64, device="cpu") + load = expert_load.to(dtype=torch.float64, device="cpu") placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank if current_placement is not None: @@ -189,12 +183,6 @@ def plan_redundant_experts( remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) layer_indices = torch.arange(num_layers, dtype=torch.int64) expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - rank_nodes = ( - torch.arange(num_ranks, dtype=torch.int64) // legacy_node_world_size - if legacy_node_world_size is not None and legacy_node_world_size < num_ranks - else None - ) - # Every iteration fills one slot per layer. Candidate expert evaluation # is vectorized across all layers and logical experts, which keeps large # GLM/Qwen planning comfortably on the CPU fast path. @@ -207,15 +195,6 @@ def plan_redundant_experts( if remaining_slots[layer, target_rank] == 0: continue candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] - # Legacy 2D/3D callers have no source-node axis. Retain the - # previous topology preference for that compatibility path; - # node-aware [S,L,N,E] planning uses only the exact load - # objective below. - if rank_nodes is not None: - existing_on_target_node = locations[layer, :, rank_nodes == rank_nodes[target_rank]].any(dim=1) - new_node_legal = candidate_legal & ~existing_on_target_node - if torch.any(new_node_legal): - candidate_legal = new_node_legal if torch.any(candidate_legal): target_ranks[layer] = target_rank legal[layer] = candidate_legal @@ -397,18 +376,13 @@ def _estimate_rank_load( expert_alignment: int | None = None, node_world_size: int | None = None, ) -> torch.Tensor: - """Estimate runtime source-node-local routing load per physical expert. + """Estimate [samples, layers, ranks] load from source-node-local routing. - ``expert_load`` accepts the historic ``[layers, experts]`` and - ``[samples, layers, experts]`` forms, which are both one source node, and - the distributed ``[samples, layers, source_nodes, experts]`` form. Source - loads are kept separate until they are assigned to physical replicas, then - combined before applying the per-expert alignment used by DeepEP. + Source loads remain separate until assigned to physical replicas, then + combine before the per-expert alignment used by DeepEP. """ - source_load, squeeze_sample, node_world_size = _as_source_node_load( - expert_load, redundant_expert_ids.shape[1], node_world_size - ) - num_samples, num_layers, num_nodes, num_logical_experts = source_load.shape + node_world_size = _resolve_node_world_size(expert_load, redundant_expert_ids.shape[1], node_world_size) + num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] assert num_logical_experts % num_ranks == 0 @@ -416,37 +390,25 @@ def _estimate_rank_load( assert expert_alignment > 0 rank_load = _expert_rank_load_all( - source_load, + expert_load, _expert_locations(redundant_expert_ids, num_logical_experts), num_nodes, node_world_size, expert_alignment, ).sum(dim=2) - return rank_load.squeeze(0) if squeeze_sample else rank_load - - -def _as_source_node_load( - expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None -) -> Tuple[torch.Tensor, bool, int]: - """Normalize load to ``[samples, layers, source_nodes, experts]``.""" - assert expert_load.ndim in (2, 3, 4) - squeeze_sample = expert_load.ndim == 2 - if expert_load.ndim == 2: - source_load = expert_load.unsqueeze(0).unsqueeze(2) - elif expert_load.ndim == 3: - source_load = expert_load.unsqueeze(2) - else: - source_load = expert_load - num_nodes = source_load.shape[2] - # Historic 2D/3D loads represent one source node containing every rank. - if expert_load.ndim < 4: - return source_load, squeeze_sample, num_ranks + return rank_load + + +def _resolve_node_world_size(expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None) -> int: + """Validate production [samples, layers, source_nodes, experts] planner loads.""" + assert expert_load.ndim == 4 + num_nodes = expert_load.shape[2] if node_world_size is None: assert num_ranks % num_nodes == 0 node_world_size = num_ranks // num_nodes assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 assert num_nodes == num_ranks // node_world_size - return source_load, squeeze_sample, node_world_size + return node_world_size def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 3f7c5fe6b6..60bcb3a39d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -207,10 +207,11 @@ def experts_with_topk( infer_state=None, clamp_limit: Optional[float] = None, alloc_tensor_func=torch.empty, + logical_topk_ids: Optional[torch.Tensor] = None, ) -> torch.Tensor: moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) if moe_capture_callback is not None: - moe_capture_callback(topk_ids) + moe_capture_callback(logical_topk_ids if logical_topk_ids is not None else topk_ids) return self.fuse_moe_impl._fused_experts( input_tensor=input_tensor, w13=self.w13, diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 627c09245d..8678cbe9f2 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -292,7 +292,6 @@ def _select_experts( group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( - hidden_states=input_tensor, gating_output=router_logits, correction_bias=correction_bias, topk=top_k, @@ -380,11 +379,6 @@ def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: if weight_pack.weight_scale is not None else None ), - weight_zero_point=( - getattr(weight_pack, "weight_zero_point", None)[:num_primary_experts_per_rank] - if getattr(weight_pack, "weight_zero_point", None) is not None - else None - ), ) cache[cache_key] = primary return primary diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index 8544b22e32..7651282b2f 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -440,7 +440,6 @@ def triton_grouped_topk( def triton_grouped_topk_eplb( - hidden_states: torch.Tensor, gating_output: torch.Tensor, correction_bias: torch.Tensor, topk: int, diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py index 441f92ca3e..63e3dea069 100644 --- a/lightllm/common/eplb_utils.py +++ b/lightllm/common/eplb_utils.py @@ -8,7 +8,7 @@ def extract_eplb_expert_tensors(weight): result = [] for pack_name in ("w13", "w2"): pack = getattr(weight, pack_name) - for value_name in ("weight", "weight_scale", "weight_zero_point"): + for value_name in ("weight", "weight_scale"): tensor = getattr(pack, value_name, None) if tensor is not None: assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 7e7138b759..a4c795a007 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -421,6 +421,7 @@ def __init__( swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, + memory_reservations=None, ): assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" assert head_dim == self.mla_head_dim, f"DeepSeek-V4 packed KV 期望 head_dim={self.mla_head_dim}" @@ -459,7 +460,16 @@ def __init__( self.layer_to_c128_idx[lid] = c128 c128 += 1 - super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) + super().__init__( + size, + dtype, + head_num, + head_dim, + layer_num, + always_copy, + mem_fraction, + memory_reservations=memory_reservations, + ) # ------------------------------------------------------------------ sizing def _planned_swa_size(self, full_size: int) -> int: diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index eb96389801..d3bfb5f3d7 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -4,6 +4,7 @@ from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd +from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args @@ -161,8 +162,7 @@ def overlap_tpsp_context_forward( x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) - weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) - + weights0, indices0, _ = self._select_experts(logits0, infer_state, layer_weight) indices0 = indices0.to(torch.long) qinput0 = experts.quantize_dispatch_input(x0) from deep_ep import ElasticBuffer @@ -175,7 +175,7 @@ def overlap_tpsp_context_forward( x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) - weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) + weights1, indices1, _ = self._select_experts(logits1, infer_state1, layer_weight) recv_x0, recv_indices0, recv_weights0, recv_count0, handle0, dispatch_hook0 = experts.dispatch( qinput0, @@ -266,7 +266,7 @@ def overlap_tpsp_token_forward( x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) - weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) + weights0, indices0, _ = self._select_experts(logits0, infer_state, layer_weight) infer_state1.call_overlap_hook() shared0 = self._ffn_tp(x0, infer_state, layer_weight) @@ -279,7 +279,7 @@ def overlap_tpsp_token_forward( x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) - weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) + weights1, indices1, _ = self._select_experts(logits1, infer_state1, layer_weight) dispatch_hook0() shared1 = self._ffn_tp(x1, infer_state1, layer_weight) @@ -523,6 +523,7 @@ def _routed_experts( x, weights, indices, + logical_topk_ids, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, ): @@ -530,6 +531,7 @@ def _routed_experts( input_tensor=x, topk_weights=weights, topk_ids=indices, + logical_topk_ids=logical_topk_ids, is_prefill=infer_state.is_prefill, infer_state=infer_state, clamp_limit=float(self.swiglu_limit), @@ -552,20 +554,25 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV x = self._tpsp_allgather(input=x, infer_state=infer_state) logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) - weights, indices = self._select_experts(logits, infer_state, layer_weight) + need_logical_ids = get_moe_capture_callback(infer_state, self.layer_num_) is not None + weights, indices, logical_topk_ids = self._select_experts( + logits, infer_state, layer_weight, return_logical_ids=need_logical_ids + ) # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) - routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) - if self.enable_ep_moe: - return routed + shared + routed = self._routed_experts(x, weights, indices, logical_topk_ids, infer_state, layer_weight) out = routed + shared - return self._tpsp_reduce(input=out, infer_state=infer_state) + return out if self.enable_ep_moe else self._tpsp_reduce(input=out, infer_state=infer_state) def _select_experts( - self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight + self, + logits, + infer_state: DeepseekV4InferStateInfo, + layer_weight: DeepseekV4TransformerLayerWeight, + return_logical_ids: bool = False, ): M = logits.shape[0] bias = None @@ -586,6 +593,32 @@ def _select_experts( if input_tokens is None: input_tokens = infer_state.input_ids + eplb = None + if layer_weight.experts_.expert_parallel_state is not None: + eplb = layer_weight.experts_.expert_parallel_state.eplb + if infer_state.is_prefill is True and eplb is not None: + if bias_vl is not None: + raise RuntimeError("DeepSeek-V4 EPLB does not support vision routing yet") + from lightllm.models.deepseek_v4.triton_kernel.moe_topk import ( + deepseek_v4_eplb_topk, + ) + + return deepseek_v4_eplb_topk( + logits=logits, + bias=bias, + input_tokens=input_tokens, + hash_indices_table=hash_indices_table, + topk=self.num_experts_per_tok, + routed_scaling_factor=self.routed_scaling_factor, + logical_to_physical_map=eplb.logical_to_physical_map, + logical_replica_count=eplb.logical_replica_count, + expert_counter=eplb.route_counter, + sample_index=eplb.next_sample_index(), + record_load=eplb.recording, + alloc_tensor_func=self.alloc_tensor, + return_logical_ids=return_logical_ids, + ) + weights = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.float32, device=logits.device) indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=indices_dtype, device=logits.device) topk_softplus_sqrt( @@ -599,7 +632,7 @@ def _select_experts( bias_vl, image_token_start, ) - return weights, indices + return weights, indices, indices if return_logical_ids else None class CompressorInfer(BaseLayerInfer): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index e18e28c17e..31afb8ed1f 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -46,6 +46,10 @@ ) from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager +from lightllm.common.eplb_utils import ( + EPLB_MAX_STAGING_DEPTH, + extract_eplb_expert_tensors, +) logger = init_logger(__name__) @@ -97,6 +101,7 @@ def _get_compress_rates(self, layer_num): def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() state_mtp_step = 0 if self.args.run_mode == "prefill" else self.args.mtp_step + reservations = self._get_post_profile_memory_reservations() self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, @@ -113,10 +118,47 @@ def _init_mem_manager(self): else self.args.cpu_cache_token_page_size ), mem_fraction=self.mem_fraction, + memory_reservations=reservations, ) self.req_manager.mem_manager = self.mem_manager return + def _get_post_profile_memory_reservations(self): + """Only buffers which are not yet visible to cuda.mem_get_info().""" + weights = self._get_eplb_weights() + staging = _get_eplb_staging_nbytes(weights) + sampling = _get_eplb_sampling_peak_nbytes(weights) + return {name: value for name, value in (("eplb_staging", staging), ("eplb_sampling", sampling)) if value} + + def _get_eplb_weights(self): + if self.is_mtp_draft_model or not self.args.enable_prefill_eplb: + return [] + weights = [] + seen = set() + for layer_weight in self.trans_layers_weight: + experts = getattr(layer_weight, "experts_", None) + state = getattr(experts, "expert_parallel_state", None) + if getattr(state, "eplb", None) is None or id(experts) in seen: + continue + seen.add(id(experts)) + weights.append(experts) + return weights + + def get_mtp_profile_weight_exclusion(self): + """Rows present only in target EPLB; the DSpark draft disables EPLB.""" + total = 0 + seen = set() + for experts in self._get_eplb_weights(): + eplb = experts.expert_parallel_state.eplb + redundant = eplb.num_redundant_experts_per_rank + for _, tensor in extract_eplb_expert_tensors(experts): + key = (tensor.data_ptr(), tensor.numel(), tensor.element_size()) + if key in seen: + continue + seen.add(key) + total += redundant * tensor[0].numel() * tensor.element_size() + return total + def _init_att_backend(self): args = get_env_start_args() if args.llm_kv_type == "None": @@ -545,3 +587,46 @@ def apply_chat_template( if tokenize: return self.tokenizer.encode(prompt, add_special_tokens=False) return prompt + + +def _get_eplb_staging_nbytes(weights) -> int: + """Owned bytes for NIXL's reusable staging rows, excluding live expert weights.""" + if not weights: + return 0 + depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) + redundant = weights[0].expert_parallel_state.eplb.num_redundant_experts_per_rank + one_row_nbytes = sum( + tensor[0].numel() * tensor.element_size() for _, tensor in extract_eplb_expert_tensors(weights[0]) + ) + return depth * redundant * one_row_nbytes + + +def _get_eplb_sampling_peak_nbytes(weights) -> int: + """Peak temporary bytes of EPLB sample collection, excluding route counters. + + _collect_local_samples keeps stack(counters), index_select output and index + temporaries live together. Counters are persistent state and intentionally + excluded. CUDA allocator rounding is represented by 512-byte alignment. + """ + counters = [] + seen = set() + for weight in weights: + state = getattr(weight, "expert_parallel_state", None) + eplb = getattr(state, "eplb", None) + counter = getattr(eplb, "route_counter", None) + if counter is None or id(counter) in seen: + continue + seen.add(id(counter)) + counters.append(counter) + if not counters: + return 0 + first = counters[0] + if any(tuple(counter.shape) != tuple(first.shape) or counter.dtype != first.dtype for counter in counters): + raise ValueError("EPLB route counters must have identical shape and dtype") + align = lambda value: (value + 511) // 512 * 512 + stack_bytes = align(len(counters) * first.numel() * first.element_size()) + # index_select has the stack shape. The int64 selection index is bounded + # by one ring axis; this is the selection peak, not a claim that every + # arithmetic intermediate remains live. + index_temp_bytes = align(first.shape[0] * 8) + return 2 * stack_bytes + index_temp_bytes diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu new file mode 100644 index 0000000000..dc9b53a4b2 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu @@ -0,0 +1,219 @@ +// Copyright 2026 LightLLM Team +// SPDX-License-Identifier: Apache-2.0 +// +// DeepSeek-V4 EPLB top-k. The warp-per-row selection follows the public +// Apache-2.0 vLLM topk_softplus_sqrt CUDA implementation, specialized to the +// DSV4 fp32, 256-expert route. + +#include +#include +#include +#include +#include +#include +#include + +namespace { +constexpr int kExperts = 256; +constexpr int kWarps = 4; +constexpr unsigned kMask = 0xffffffffu; + +__device__ __forceinline__ float score(float x) { + return sqrtf(fmaxf(x, 0.f) + log1pf(expf(-fabsf(x)))); +} + +__device__ __forceinline__ bool better(float value, int id, float other_value, int other_id) { + return value > other_value || (value == other_value && id < other_id); +} + +template +__global__ void moe_topk_eplb_kernel( + const float* __restrict__ logits, const float* __restrict__ bias, + const int64_t* __restrict__ input_ids, const int64_t* __restrict__ tid2eid, + float* __restrict__ weights, int64_t* __restrict__ physical_ids, + int64_t* __restrict__ logical_ids, const int32_t* __restrict__ logical_to_physical, + const int32_t* __restrict__ replica_count, int64_t* __restrict__ counter, + int64_t sample_index, int tokens, int topk, int map_slots, int counter_experts, + float routed_scaling_factor) { + const int lane = threadIdx.x & 31; + const int token = blockIdx.x * kWarps + threadIdx.x / 32; + if (token >= tokens) return; + const float* row = logits + token * kExperts; + int selected[8]; + float selected_score[8]; + #pragma unroll + for (int i = 0; i < 8; ++i) { + selected[i] = -1; + selected_score[i] = 0.f; + } + + if constexpr (IsHash) { + if (lane == 0) { + const int64_t input = input_ids[token]; + #pragma unroll + for (int k = 0; k < 8; ++k) { + if (k >= topk) break; + const int id = static_cast(tid2eid[input * topk + k]); + selected[k] = id; + selected_score[k] = score(row[id]); + } + } + } else { + float local_choice[8]; + #pragma unroll + for (int j = 0; j < 8; ++j) { + const int id = lane + j * 32; + local_choice[j] = score(row[id]) + bias[id]; + } + #pragma unroll + for (int k = 0; k < 8; ++k) { + if (k >= topk) break; + float best = -INFINITY; + int best_id = kExperts; + #pragma unroll + for (int j = 0; j < 8; ++j) { + const int id = lane + j * 32; + bool used = false; + #pragma unroll + for (int prev = 0; prev < 8; ++prev) used |= prev < k && selected[prev] == id; + if (!used && better(local_choice[j], id, best, best_id)) { best = local_choice[j]; best_id = id; } + } + for (int offset = 16; offset > 0; offset >>= 1) { + const float other = __shfl_down_sync(kMask, best, offset); + const int other_id = __shfl_down_sync(kMask, best_id, offset); + if (lane + offset < 32 && better(other, other_id, best, best_id)) { best = other; best_id = other_id; } + } + best_id = __shfl_sync(kMask, best_id, 0); + selected[k] = best_id; + selected_score[k] = score(row[best_id]); + } + } + if (lane == 0) { + float sum = 0.f; + #pragma unroll + for (int k = 0; k < 8; ++k) { + if (k < topk) sum += selected_score[k]; + } + #pragma unroll + for (int k = 0; k < 8; ++k) { + if (k >= topk) break; + const int id = selected[k]; + uint32_t replica = 0; + if constexpr (!SingleToken) { + const uint32_t token_hash = static_cast(token) * 2654435769u; + const uint32_t expert_hash = static_cast(id) * 2246822519u; + replica = (token_hash + expert_hash) % static_cast(replica_count[id]); + } + weights[token * topk + k] = selected_score[k] / fmaxf(sum, 1e-20f) * routed_scaling_factor; + physical_ids[token * topk + k] = logical_to_physical[id * map_slots + replica]; + if constexpr (ReturnLogical) logical_ids[token * topk + k] = id; + if constexpr (RecordLoad) { + atomicAdd( + reinterpret_cast(counter + sample_index * counter_experts + id), + 1ULL); + } + } + } +} + +template +void launch(const torch::Tensor& logits, const torch::Tensor& bias, const torch::Tensor& input_ids, + const torch::Tensor& tid2eid, const torch::Tensor& weights, const torch::Tensor& physical_ids, + const torch::Tensor& logical_ids, const torch::Tensor& logical_to_physical, + const torch::Tensor& replica_count, const torch::Tensor& counter, int64_t sample_index, + float routed_scaling_factor) { + const int tokens = logits.size(0), topk = weights.size(1); + moe_topk_eplb_kernel + <<< (tokens + kWarps - 1) / kWarps, kWarps * 32, 0, at::cuda::getCurrentCUDAStream() >>>( + logits.data_ptr(), bias.data_ptr(), input_ids.data_ptr(), tid2eid.data_ptr(), + weights.data_ptr(), physical_ids.data_ptr(), logical_ids.data_ptr(), + logical_to_physical.data_ptr(), replica_count.data_ptr(), counter.data_ptr(), + sample_index, tokens, topk, logical_to_physical.size(1), counter.size(1), routed_scaling_factor); + C10_CUDA_KERNEL_LAUNCH_CHECK(); +} + +template +void dispatch(const torch::Tensor& logits, const torch::Tensor& bias, const torch::Tensor& input_ids, + const torch::Tensor& tid2eid, const torch::Tensor& weights, const torch::Tensor& physical_ids, + const torch::Tensor& logical_ids, const torch::Tensor& logical_to_physical, + const torch::Tensor& replica_count, const torch::Tensor& counter, int64_t sample_index, + bool is_hash, bool single_token, float routed_scaling_factor) { + if (is_hash && single_token) { + launch(logits, bias, input_ids, tid2eid, weights, physical_ids, + logical_ids, logical_to_physical, replica_count, counter, + sample_index, routed_scaling_factor); + } else if (is_hash) { + launch(logits, bias, input_ids, tid2eid, weights, physical_ids, + logical_ids, logical_to_physical, replica_count, counter, + sample_index, routed_scaling_factor); + } else if (single_token) { + launch(logits, bias, input_ids, tid2eid, weights, physical_ids, + logical_ids, logical_to_physical, replica_count, counter, + sample_index, routed_scaling_factor); + } else { + launch(logits, bias, input_ids, tid2eid, weights, physical_ids, + logical_ids, logical_to_physical, replica_count, counter, + sample_index, routed_scaling_factor); + } +} +} // namespace + +void moe_topk_eplb(torch::Tensor logits, torch::Tensor bias, torch::Tensor input_ids, torch::Tensor tid2eid, + torch::Tensor weights, torch::Tensor physical_ids, torch::Tensor logical_ids, + torch::Tensor logical_to_physical, torch::Tensor replica_count, torch::Tensor counter, + int64_t sample_index, bool record_load, bool return_logical_ids, bool is_hash, + float routed_scaling_factor) { + const auto check = [&](const torch::Tensor& tensor, const char* name, torch::ScalarType dtype) { + TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA"); + TORCH_CHECK(tensor.device() == logits.device(), name, " must share logits device"); + TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous"); + TORCH_CHECK(tensor.scalar_type() == dtype, name, " has unexpected dtype"); + }; + check(logits, "logits", torch::kFloat); + check(bias, "bias", torch::kFloat); + check(input_ids, "input_ids", torch::kLong); + check(tid2eid, "tid2eid", torch::kLong); + check(weights, "weights", torch::kFloat); + check(physical_ids, "physical_ids", torch::kLong); + check(logical_ids, "logical_ids", torch::kLong); + check(logical_to_physical, "logical_to_physical", torch::kInt); + check(replica_count, "replica_count", torch::kInt); + check(counter, "counter", torch::kLong); + TORCH_CHECK(logits.dim() == 2 && logits.size(1) == kExperts, "logits must be [M, 256]"); + const int64_t tokens = logits.size(0); + const int64_t topk = weights.size(1); + TORCH_CHECK(topk >= 1 && topk <= 8, "topk must be in [1, 8]"); + TORCH_CHECK(weights.dim() == 2 && weights.size(0) == tokens, "weights must be [M, K]"); + TORCH_CHECK(physical_ids.sizes() == weights.sizes(), "physical_ids must be [M, K]"); + TORCH_CHECK(logical_ids.sizes() == weights.sizes(), "logical_ids must be [M, K]"); + TORCH_CHECK(logical_to_physical.dim() == 2 && logical_to_physical.size(0) == kExperts && + logical_to_physical.size(1) > 0, + "logical_to_physical must be [256, map_slots]"); + TORCH_CHECK(replica_count.dim() == 1 && replica_count.size(0) == kExperts, "replica_count must be [256]"); + TORCH_CHECK(counter.dim() == 2 && counter.size(0) > 0 && counter.size(1) == kExperts, + "counter must be [rows, 256]"); + TORCH_CHECK(sample_index >= 0 && sample_index < counter.size(0), "sample_index outside counter rows"); + c10::cuda::CUDAGuard guard(logits.device()); + const bool single = tokens == 1; + if (is_hash) { + TORCH_CHECK(input_ids.dim() == 1 && input_ids.size(0) == tokens, "input_ids must be [M]"); + TORCH_CHECK(tid2eid.dim() == 2 && tid2eid.size(1) == topk, "tid2eid must be [vocab, K]"); + } else { + TORCH_CHECK(bias.dim() == 1 && bias.size(0) == kExperts, "bias must be [256]"); + } + if (record_load && return_logical_ids) { + dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, + logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); + } else if (record_load) { + dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, + logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); + } else if (return_logical_ids) { + dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, + logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); + } else { + dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, + logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); + } +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("moe_topk_eplb", &moe_topk_eplb); } diff --git a/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py b/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py new file mode 100644 index 0000000000..800055e9a4 --- /dev/null +++ b/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py @@ -0,0 +1,81 @@ +import functools +import hashlib +import os + +import torch + + +@torch.no_grad() +def deepseek_v4_eplb_topk( + logits: torch.Tensor, + bias: torch.Tensor | None, + input_tokens: torch.Tensor | None, + hash_indices_table: torch.Tensor | None, + topk: int, + routed_scaling_factor: float, + logical_to_physical_map: torch.Tensor, + logical_replica_count: torch.Tensor, + expert_counter: torch.Tensor, + sample_index: int, + record_load: bool, + alloc_tensor_func=torch.empty, + return_logical_ids: bool = False, +): + """Select DeepSeek-V4 routes and emit EPLB physical expert IDs.""" + token_num, num_experts = logits.shape + if num_experts != 256: + raise RuntimeError(f"DeepSeek-V4 EPLB fused top-k requires 256 experts, got {num_experts}") + if hash_indices_table is not None: + topk = hash_indices_table.shape[1] + weights = alloc_tensor_func((token_num, topk), dtype=torch.float32, device=logits.device) + physical_ids = alloc_tensor_func((token_num, topk), dtype=torch.long, device=logits.device) + logical_ids = ( + alloc_tensor_func((token_num, topk), dtype=torch.long, device=logits.device) if return_logical_ids else None + ) + if token_num == 0: + return weights, physical_ids, logical_ids + _load_cuda().moe_topk_eplb( + logits.contiguous(), + bias.contiguous() if bias is not None else logits, + input_tokens.contiguous() if input_tokens is not None else physical_ids, + hash_indices_table.contiguous() if hash_indices_table is not None else physical_ids, + weights, + physical_ids, + logical_ids if logical_ids is not None else physical_ids, + logical_to_physical_map, + logical_replica_count, + expert_counter, + sample_index, + record_load, + return_logical_ids, + hash_indices_table is not None, + routed_scaling_factor, + ) + return weights, physical_ids, logical_ids + + +@functools.lru_cache(maxsize=1) +def _load_cuda(): + from torch.utils.cpp_extension import load + + source_path = os.path.join(os.path.dirname(__file__), "csrc", "moe_topk_eplb.cu") + flags = ["-O3"] + with open(source_path, "rb") as source_file: + source = source_file.read() + capability = torch.cuda.get_device_capability() + cache_key = b"\0".join( + [ + source, + " ".join(flags).encode(), + torch.__version__.encode(), + str(torch.version.cuda).encode(), + f"sm{capability[0]}{capability[1]}".encode(), + os.environ.get("TORCH_CUDA_ARCH_LIST", "").encode(), + ] + ) + return load( + name=f"lightllm_dsv4_eplb_topk_v1_{hashlib.sha256(cache_key).hexdigest()[:16]}", + sources=[source_path], + extra_cuda_cflags=flags, + verbose=False, + ) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py index a73f9a54b7..456565e81d 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py @@ -524,14 +524,10 @@ def _start_rebalance(self, result): def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: - if rank_load.ndim == 2: - critical = rank_load.max(dim=1).values - mean = rank_load.mean(dim=1) - elif rank_load.ndim == 3: - critical = rank_load.max(dim=2).values.sum(dim=0) - mean = rank_load.mean(dim=2).sum(dim=0) - else: - raise ValueError("rank_load must be [layers, ranks] or [samples, layers, ranks]") + if rank_load.ndim != 3: + raise ValueError("rank_load must be [samples, layers, ranks]") + critical = rank_load.max(dim=2).values.sum(dim=0) + mean = rank_load.mean(dim=2).sum(dim=0) layer_imbalance = critical / mean.clamp_min(1.0) sorted_imbalance = torch.sort(layer_imbalance).values p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py index 9d0c9f120e..e5737fcb7a 100644 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py @@ -99,8 +99,6 @@ def build_transfer_plan( class _EPLBTransferBase: """Shared live/staging buffers and publish/commit lifecycle.""" - staging_depth = 1 - def __init__(self, weights, transfer_group, global_rank, world_size): self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] self.transfer_group = transfer_group @@ -109,7 +107,7 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank self.device = weights[0].w13.weight.device self.live = [extract_eplb_expert_tensors(weight) for weight in weights] - self._validate_live_layout(weights) + self._validate_live_layout() num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank self.staging = [ [ @@ -137,7 +135,7 @@ def __init__(self, weights, transfer_group, global_rank, world_size): self._thread = None self._needs_staging_reuse_barrier = False - def _validate_live_layout(self, weights) -> None: + def _validate_live_layout(self) -> None: reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): @@ -147,9 +145,6 @@ def _validate_live_layout(self, weights) -> None: state.num_redundant_experts_per_rank == num_redundant_slots_per_rank ), "EPLB redundant slot count must match" - def _copy_batch(self, batch, prepared_batch) -> None: - raise NotImplementedError - def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): return [ [ @@ -161,20 +156,9 @@ def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]] for batch_start in range(0, len(layer_plans), self.staging_depth) ] - def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [(batch, None) for batch in self._make_batches(layer_plans)] - - def _start_transfer_generation(self) -> None: - """Prepare backend state after the in-flight worker check succeeds.""" - - def _finish_transfer_generation(self) -> None: - """Release backend state only after the migration worker has joined.""" - - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches=None) -> None: + def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches) -> None: if self._thread is not None and self._thread.is_alive(): raise RuntimeError("EPLB transfer is already in flight") - if prepared_batches is None: - prepared_batches = self.prepare_transfer(layer_plans) expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth if len(prepared_batches) != expected_batch_count: raise ValueError("EPLB prepared batch count does not match layer-plan batches") diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index b9b10923b6..5ab0ee8865 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -171,6 +171,7 @@ def get_eplb_placement_stickiness() -> float: return value +@lru_cache(maxsize=None) def get_triton_autotune_level(): return int(os.getenv("LIGHTLLM_TRITON_AUTOTUNE_LEVEL", 0)) diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py index bb3e97fb8e..fd656e07c3 100644 --- a/unit_tests/common/fused_moe/test_eplb.py +++ b/unit_tests/common/fused_moe/test_eplb.py @@ -303,11 +303,15 @@ def test_build_initial_redundant_expert_ids( def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): - expert_load = torch.tensor( - [ - [100, 90, 80, 70, 60, 50, 40, 30], - [30, 40, 50, 60, 70, 80, 90, 100], - ] + expert_load = ( + torch.tensor( + [ + [100, 90, 80, 70, 60, 50, 40, 30], + [30, 40, 50, 60, 70, 80, 90, 100], + ] + ) + .unsqueeze(0) + .unsqueeze(2) ) placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) @@ -318,7 +322,7 @@ def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): - samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]) + samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]).unsqueeze(2) placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] @@ -330,7 +334,7 @@ def critical(candidate): def test_select_improving_placements_rejects_regressing_layer(): - expert_load = torch.tensor([[8649, 5740, 5002, 3441]]) + expert_load = torch.tensor([[8649, 5740, 5002, 3441]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) regressing_candidate = torch.tensor([[[1], [0]]]) @@ -350,7 +354,7 @@ def test_select_improving_placements_rejects_regressing_layer(): def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): - expert_load = torch.tensor([[1, 2, 1, 17]]) + expert_load = torch.tensor([[1, 2, 1, 17]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[3], [0]]]) candidate = torch.tensor([[[3], [1]]]) @@ -369,7 +373,7 @@ def test_select_improving_placements_rejects_near_balance_when_gain_is_below_thr def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): - expert_load = torch.tensor([[100, 129, 100, 129]]) + expert_load = torch.tensor([[100, 129, 100, 129]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [1]]]) @@ -388,7 +392,7 @@ def test_select_improving_placements_accepts_alignment_aware_gain_even_when_curr def test_select_improving_placements_rejects_insufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 6, 7]]) + expert_load = torch.tensor([[1, 1, 6, 7]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -422,7 +426,7 @@ def test_select_improving_placements_rejects_insufficient_rebalance_gain(): def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( rebalance_gain_threshold, ): - expert_load = torch.tensor([[1, 1, 1, 2]]) + expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -436,7 +440,7 @@ def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( def test_select_improving_placements_accepts_sufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 1, 2]]) + expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[3], [0]]]) @@ -457,7 +461,7 @@ def test_select_improving_placements_accepts_sufficient_rebalance_gain(): def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): - expert_load = torch.tensor([[1, 1, 1, 8]]) + expert_load = torch.tensor([[1, 1, 1, 8]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]]]) raw_improving_candidate = torch.tensor([[[3], [0]]]) @@ -478,15 +482,15 @@ def test_select_improving_placements_rejects_raw_improvement_that_does_not_impro def test_estimate_rank_load_aligns_each_sample_before_accumulation(): - samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]) + samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]).unsqueeze(2) placement = torch.tensor([[[2], [0]]]) per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) - accumulated = _estimate_rank_load(samples.sum(dim=0), placement, expert_alignment=128) + accumulated = _estimate_rank_load(samples.sum(dim=0, keepdim=True), placement, expert_alignment=128) assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) - assert torch.equal(accumulated[0], torch.tensor([128.0, 128.0])) + assert torch.equal(accumulated[0, 0], torch.tensor([128.0, 128.0])) def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): @@ -496,7 +500,7 @@ def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchan [[172, 278, 51, 238]], [[249, 291, 284, 183]], ] - ) + ).unsqueeze(2) current = torch.tensor([[[2], [0]]]) mean_inflating_candidate = torch.tensor([[[2], [1]]]) @@ -525,7 +529,7 @@ def test_select_improving_placements_accepts_five_percent_critical_reduction(): [[287, 175, 236, 179]], [[316, 99, 266, 353]], ] - ) + ).unsqueeze(2) current = torch.tensor([[[2], [0]]]) candidate = torch.tensor([[[2], [1]]]) @@ -540,7 +544,7 @@ def test_select_improving_placements_accepts_five_percent_critical_reduction(): def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): # Layer 0 becomes better, but layer 1 dominates model critical load. The # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]) + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]], [[2], [0]]]) candidate = torch.tensor([[[3], [0]], [[2], [0]]]) @@ -555,7 +559,7 @@ def test_select_improving_placements_rejects_single_layer_gain_below_model_thres def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]) + expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]).unsqueeze(0).unsqueeze(2) current = torch.tensor([[[2], [0]], [[2], [0]]]) candidate = torch.tensor([[[3], [0]], [[2], [0]]]) @@ -674,22 +678,23 @@ def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, assert torch.all(maps_by_layer[~valid] == -1) -def test_plan_redundant_experts_prefers_first_replica_on_new_node(): - # One redundant slot per rank leaves legal alternatives on both nodes; - # topology preference therefore puts every first replica away from its - # primary node before considering same-node duplicates. +def test_plan_redundant_experts_prefers_local_node_load_relief(): + source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) + source_load[0, 0, 0, 0] = 1024 placement = plan_redundant_experts( - torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]), + source_load, num_ranks=4, num_redundant_experts_per_rank=1, + expert_alignment=128, node_world_size=2, ) - for rank, expert in enumerate(placement[0, :, 0].tolist()): - assert expert // 2 // 2 != rank // 2 + assert placement[0, 1, 0] == 0 + predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) + assert torch.equal(predicted, _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128)) def test_plan_redundant_experts_single_node_matches_default_behavior(): - load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]) + load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]).unsqueeze(0).unsqueeze(2) default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) assert torch.equal(single_node, default) @@ -705,7 +710,7 @@ def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) - collapsed_global = _estimate_rank_load(source_load.sum(dim=2), placement, expert_alignment=128) + collapsed_global = _estimate_rank_load(source_load.sum(dim=2, keepdim=True), placement, expert_alignment=128) assert torch.equal(predicted, runtime) assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) @@ -783,7 +788,7 @@ def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: def test_sticky_plan_reproduces_current_when_load_unchanged(): generator = torch.Generator().manual_seed(7) - load = torch.randint(1, 1000, (3, 16, 32), generator=generator) + load = torch.randint(1, 1000, (3, 16, 32), generator=generator).unsqueeze(2) placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) replanned = plan_redundant_experts( @@ -801,9 +806,9 @@ def test_sticky_plan_reproduces_current_when_load_unchanged(): def test_sticky_plan_bounded_moves_under_small_perturbation(): generator = torch.Generator().manual_seed(11) - load = torch.randint(100, 1000, (4, 16, 32), generator=generator) + load = torch.randint(100, 1000, (4, 16, 32), generator=generator).unsqueeze(2) placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - noise = torch.rand((4, 16, 32), generator=generator) * 0.1 + 0.95 + noise = torch.rand((4, 16, 32), generator=generator).unsqueeze(2) * 0.1 + 0.95 perturbed = (load.double() * noise).round().to(torch.int64) sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) @@ -828,6 +833,8 @@ def test_sticky_plan_still_churns_under_phase_shift(): for layer in range(layers): before[layer, (4 * layer + offsets) % experts] = 5000 after[layer, (4 * layer + 16 + offsets) % experts] = 5000 + before = before.unsqueeze(0).unsqueeze(2) + after = after.unsqueeze(0).unsqueeze(2) placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) replanned = plan_redundant_experts( @@ -902,7 +909,7 @@ def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): broadcasts = [] def fixed_selector(*_args, **_kwargs): - rank_load = torch.full((1, 4), 100.0) + rank_load = torch.full((1, 1, 4), 100.0) return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load def record_broadcast(result_list, **_kwargs): @@ -922,10 +929,10 @@ def record_broadcast(result_list, **_kwargs): assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) -def test_stickiness_zero_matches_legacy(): +def test_stickiness_zero_matches_unbiased_plan(): generator = torch.Generator().manual_seed(17) - load = torch.randint(1, 1000, (2, 8, 16), generator=generator) - legacy = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) + load = torch.randint(1, 1000, (2, 8, 16), generator=generator).unsqueeze(2) + unbiased = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() replanned = plan_redundant_experts( @@ -936,7 +943,7 @@ def test_stickiness_zero_matches_legacy(): stickiness=0.0, ) - assert torch.equal(replanned, legacy) + assert torch.equal(replanned, unbiased) def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast( @@ -2003,19 +2010,17 @@ def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): ] -def test_extract_expert_tensors_includes_weight_scale_and_zero_point_in_order(): +def test_extract_expert_tensors_includes_weight_and_scale_in_order(): class Pack: - def __init__(self, offset, scale=True, zero=True): + def __init__(self, offset, scale=True): self.weight = torch.full((3, 2), offset) self.weight_scale = torch.full((3, 1), offset + 1) if scale else None - self.weight_zero_point = torch.full((3, 1), offset + 2) if zero else None - weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False, zero=False)})() + weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False)})() tensors = extract_eplb_expert_tensors(weight) assert [name for name, _ in tensors] == [ "w13.weight", "w13.weight_scale", - "w13.weight_zero_point", "w2.weight", ] @@ -2358,6 +2363,8 @@ def synchronize(self): transfer._error = None transfer._thread = None transfer._needs_staging_reuse_barrier = True + transfer._start_transfer_generation = lambda: None + transfer._finish_transfer_generation = lambda: None transfer.transfer_group = "transfer-group" copied = [] @@ -2377,7 +2384,7 @@ def copy_batch(batch, _prepared_batch): ) plans = [(0, []), (1, []), (2, [])] - prepared_batches = transfer.prepare_transfer(plans) + prepared_batches = [(batch, None) for batch in transfer._make_batches(plans)] monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) transfer.start(plans, prepared_batches) deadline = time.monotonic() + 2 @@ -2426,6 +2433,7 @@ def make_transfer(finalize): transfer._thread = None transfer._needs_staging_reuse_barrier = False transfer._copy_batch = lambda _batch, _prepared_batch: None + transfer._start_transfer_generation = lambda: None transfer._finish_transfer_generation = finalize return transfer @@ -2433,13 +2441,13 @@ def make_transfer(finalize): finalized_before_publish = [] success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) - success.start([(0, [])]) + success.start([(0, [])], [([(0, [], 0, [])], None)]) success.finish() assert finalized_before_publish == [0] assert success.pending_layers() == [(0, 0)] failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) - failed.start([(0, [])]) + failed.start([(0, [])], [([(0, [], 0, [])], None)]) failed._thread.join() assert list(failed._pending) == [] with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: @@ -2899,12 +2907,12 @@ def enqueue(self, descriptor, stream): [ ("w13.weight", 1000, 32), ("w13.weight_scale", 2000, 32), - ("w13.weight_zero_point", 3000, 32), + ("w2.weight", 3000, 32), ] ] transfer._push_staging_row_layout = { - 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w13.weight_zero_point", 6000, 32)]], - 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w13.weight_zero_point", 9000, 32)]], + 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w2.weight", 6000, 32)]], + 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w2.weight", 9000, 32)]], } transfer._get_remote_read = lambda *_args: None transfer._wait_xfers = lambda _xfers: None @@ -3537,7 +3545,6 @@ def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), ) fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( - hidden_states, gating_output, correction_bias, topk, @@ -3607,7 +3614,6 @@ def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens) expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - hidden_states=torch.empty((tokens, 1), dtype=torch.float32, device="cuda"), gating_output=gating_output, correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), topk=topk, @@ -3638,7 +3644,6 @@ def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): experts = 64 counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - hidden_states=torch.empty((0, 1), device="cuda"), gating_output=torch.empty((0, experts), device="cuda"), correction_bias=None, topk=4, diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py index 46c218ad45..3eff70f940 100644 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py @@ -90,7 +90,7 @@ def _run_layers( callback=lambda layer_index: None, before_commit_callback=lambda layer_index: None, ): - transfer.start(layer_plans) + transfer.start(layer_plans, transfer.prepare_transfer(layer_plans)) committed = 0 while committed < len(layer_plans): pending = _wait_for_ready_prefix(transfer, control_group) diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py index 95d41d521b..bf9aaf0ead 100644 --- a/unit_tests/models/deepseek_v4/test_memory_profile.py +++ b/unit_tests/models/deepseek_v4/test_memory_profile.py @@ -3,10 +3,59 @@ import pytest import torch +from lightllm.common.eplb_utils import extract_eplb_expert_tensors from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager +from lightllm.models.deepseek_v4.model import ( + DeepseekV4TpPartModel, + _get_eplb_sampling_peak_nbytes, + _get_eplb_staging_nbytes, +) from lightllm.utils import profile_max_tokens +def _expert(rows=4, redundant=2): + def pack(cols, scale=True): + return SimpleNamespace( + weight=torch.empty((rows, cols), dtype=torch.uint8), + weight_scale=torch.empty((rows, 2), dtype=torch.float32) if scale else None, + ) + + counter = torch.zeros((5, 4), dtype=torch.int64) + return SimpleNamespace( + w13=pack(8), + w2=pack(4, False), + expert_parallel_state=SimpleNamespace( + eplb=SimpleNamespace(num_redundant_experts_per_rank=redundant, route_counter=counter) + ), + ) + + +@pytest.mark.parametrize( + "enable,draft,redundant,staging,sampling,exclusion", + [ + (True, False, 2, 40, 1536, 40), + (True, False, 0, 0, 1536, 0), + (False, False, 2, 0, 0, 0), + (True, True, 2, 0, 0, 0), + ], +) +def test_eplb_helpers_and_model_dedupe(enable, draft, redundant, staging, sampling, exclusion): + expert = _expert(redundant=redundant) + model = DeepseekV4TpPartModel.__new__(DeepseekV4TpPartModel) + model.is_mtp_draft_model = draft + model.args = SimpleNamespace(enable_prefill_eplb=enable) + model.trans_layers_weight = [SimpleNamespace(experts_=expert), SimpleNamespace(experts_=expert)] + weights = model._get_eplb_weights() + assert weights == ([expert] if enable and not draft else []) + assert _get_eplb_staging_nbytes(weights) == staging + assert _get_eplb_sampling_peak_nbytes(weights) == sampling + assert model.get_mtp_profile_weight_exclusion() == exclusion + assert sum(tensor[0].numel() * tensor.element_size() for _, tensor in extract_eplb_expert_tensors(expert)) == 20 + assert model._get_post_profile_memory_reservations() == { + name: value for name, value in (("eplb_staging", staging), ("eplb_sampling", sampling)) if value + } + + @pytest.mark.parametrize("exclusion,expected", [(0, 1000), (200, 800), (None, 1000)]) def test_mtp_profile_exclusion_adjustment(monkeypatch, exclusion, expected): seen = [] diff --git a/unit_tests/models/deepseek_v4/test_moe_topk.py b/unit_tests/models/deepseek_v4/test_moe_topk.py new file mode 100644 index 0000000000..1236684bc8 --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_moe_topk.py @@ -0,0 +1,91 @@ +import pytest +import torch +import torch.nn.functional as F + +from lightllm.models.deepseek_v4.triton_kernel.moe_topk import deepseek_v4_eplb_topk + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("is_hash", [False, True]) +@pytest.mark.parametrize("record_load", [False, True]) +@pytest.mark.parametrize("token_num", [1, 4]) +@pytest.mark.parametrize("return_logical_ids", [False, True]) +def test_deepseek_v4_eplb_topk_matches_reference(is_hash, record_load, token_num, return_logical_ids): + torch.manual_seed(0) + device = "cuda" + num_experts, topk = 256, 6 + logits = torch.randn((token_num, num_experts), device=device, dtype=torch.float32) + bias = None if is_hash else torch.randn((num_experts,), device=device, dtype=torch.float32) + input_tokens = torch.arange(token_num, device=device, dtype=torch.long) + hash_indices = torch.randint(num_experts, (token_num + 4, topk), device=device) if is_hash else None + logical_to_physical = torch.arange(num_experts * 2, device=device, dtype=torch.int32).view(num_experts, 2) + replica_count = torch.full((num_experts,), 2, device=device, dtype=torch.int32) + counter = torch.zeros((3, num_experts), device=device, dtype=torch.int64) + + weights, physical_ids, logical_ids = deepseek_v4_eplb_topk( + logits=logits, + bias=bias, + input_tokens=input_tokens if is_hash else None, + hash_indices_table=hash_indices, + topk=topk, + routed_scaling_factor=1.7, + logical_to_physical_map=logical_to_physical, + logical_replica_count=replica_count, + expert_counter=counter, + sample_index=1, + record_load=record_load, + return_logical_ids=return_logical_ids, + ) + scores = F.softplus(logits).sqrt() + expected_logical_ids = hash_indices[input_tokens] if is_hash else (scores + bias).topk(topk, dim=-1).indices + expected_weights = scores.gather(1, expected_logical_ids) + expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) * 1.7 + token_indices = torch.arange(token_num, device=device, dtype=torch.int64).view(-1, 1) + if token_num == 1: + replica_indices = torch.zeros_like(expected_logical_ids) + else: + token_phase = (token_indices * 2654435769) & 0xFFFFFFFF + expert_phase = (expected_logical_ids * 2246822519) & 0xFFFFFFFF + replica_indices = ((token_phase + expert_phase) & 0xFFFFFFFF) % 2 + expected_physical_ids = logical_to_physical[expected_logical_ids, replica_indices].to(torch.long) + torch.cuda.synchronize() + + torch.testing.assert_close(weights, expected_weights) + if return_logical_ids: + torch.testing.assert_close(logical_ids, expected_logical_ids) + else: + assert logical_ids is None + torch.testing.assert_close(physical_ids, expected_physical_ids) + expected_counter = torch.zeros_like(counter) + if record_load: + expected_counter[1].scatter_add_( + 0, + expected_logical_ids.reshape(-1), + torch.ones_like(expected_logical_ids.reshape(-1), dtype=torch.int64), + ) + torch.testing.assert_close(counter, expected_counter) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +def test_deepseek_v4_eplb_topk_empty_input(): + empty_logits = torch.empty((0, 256), device="cuda", dtype=torch.float32) + map_ = torch.arange(256, device="cuda", dtype=torch.int32).view(256, 1) + counter = torch.zeros((1, 256), device="cuda", dtype=torch.int64) + weights, physical_ids, logical_ids = deepseek_v4_eplb_topk( + logits=empty_logits, + bias=torch.zeros((256,), device="cuda"), + input_tokens=None, + hash_indices_table=None, + topk=6, + routed_scaling_factor=1.0, + logical_to_physical_map=map_, + logical_replica_count=torch.ones((256,), device="cuda", dtype=torch.int32), + expert_counter=counter, + sample_index=0, + record_load=True, + return_logical_ids=True, + ) + assert weights.shape == (0, 6) and weights.dtype == torch.float32 + assert physical_ids.shape == (0, 6) and physical_ids.dtype == torch.long + assert logical_ids.shape == (0, 6) and logical_ids.dtype == torch.long + assert counter.sum().item() == 0 diff --git a/unit_tests/models/deepseek_v4/test_vision_integration.py b/unit_tests/models/deepseek_v4/test_vision_integration.py index 5a91049ec0..80086f04d8 100644 --- a/unit_tests/models/deepseek_v4/test_vision_integration.py +++ b/unit_tests/models/deepseek_v4/test_vision_integration.py @@ -336,10 +336,11 @@ def test_router_uses_bias_vl_only_for_image_tokens(is_hash): gate_tid2eid_=SimpleNamespace(weight=hash_table), gate_bias_=SimpleNamespace(weight=text_bias), gate_bias_vl_=SimpleNamespace(weight=vision_bias), + experts_=SimpleNamespace(expert_parallel_state=None), ) infer_state = SimpleNamespace(is_prefill=True, input_ids=torch.tensor([2, 100_000], device="cuda")) - weights, indices = router._select_experts(logits, infer_state, layer_weight) + weights, indices, _ = router._select_experts(logits, infer_state, layer_weight) scores = torch.sqrt(torch.nn.functional.softplus(logits)) text_indices = hash_table[infer_state.input_ids[0]] if is_hash else (scores[0] + text_bias).topk(6).indices @@ -351,6 +352,25 @@ def test_router_uses_bias_vl_only_for_image_tokens(is_hash): torch.testing.assert_close(weights, expected, rtol=2e-5, atol=1e-6) +def test_eplb_rejects_vision_routing(): + from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer + + router = DeepseekV4TransformerLayerInfer.__new__(DeepseekV4TransformerLayerInfer) + router.has_vision = True + router.is_hash = False + router.vocab_size = 32 + logits = torch.zeros((1, 256), dtype=torch.float32) + infer_state = SimpleNamespace(is_prefill=True, input_ids=torch.tensor([100_000])) + layer_weight = SimpleNamespace( + gate_bias_=SimpleNamespace(weight=torch.zeros(256, dtype=torch.float32)), + gate_bias_vl_=SimpleNamespace(weight=torch.zeros(256, dtype=torch.float32)), + experts_=SimpleNamespace(expert_parallel_state=SimpleNamespace(eplb=object())), + ) + + with pytest.raises(RuntimeError, match="does not support vision routing"): + router._select_experts(logits, infer_state, layer_weight) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_build_image_visibility_scatter(): from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_image_visibility From 2ce90e78a8bebda27c0545e70a7a4da76f63c1a5 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 9 Sep 2026 02:59:08 +0000 Subject: [PATCH 170/214] perf(deepseek_v4): enable fused SDPA for vision attention --- lightllm/models/deepseek_v4/deepseek_v4_visual.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/lightllm/models/deepseek_v4/deepseek_v4_visual.py b/lightllm/models/deepseek_v4/deepseek_v4_visual.py index 9f0193d280..791361f911 100644 --- a/lightllm/models/deepseek_v4/deepseek_v4_visual.py +++ b/lightllm/models/deepseek_v4/deepseek_v4_visual.py @@ -77,8 +77,13 @@ def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torc q, k, v = (t.view(n, self.n_heads, self.head_dim) for t in self.wqkv(x).chunk(3, dim=-1)) q = apply_rotary(q, cos, sin) k = apply_rotary(k, cos, sin) - o = F.scaled_dot_product_attention(q.transpose(0, 1), k.transpose(0, 1), v.transpose(0, 1)) - return self.wo(o.transpose(0, 1).reshape(n, -1)) + # Keep the batch dimension so SDPA can select fused attention kernels. + o = F.scaled_dot_product_attention( + q.transpose(0, 1).unsqueeze(0), + k.transpose(0, 1).unsqueeze(0), + v.transpose(0, 1).unsqueeze(0), + ) + return self.wo(o.squeeze(0).transpose(0, 1).reshape(n, -1)) class MLP(nn.Module): From 924c62e14a0158260e88a0f52a6a05db206abfd9 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 9 Sep 2026 03:59:37 +0000 Subject: [PATCH 171/214] perf(vision): use shared ViT attention backends for DeepSeek V4 --- .../models/deepseek_v4/deepseek_v4_visual.py | 20 +++++++++---------- 1 file changed, 9 insertions(+), 11 deletions(-) diff --git a/lightllm/models/deepseek_v4/deepseek_v4_visual.py b/lightllm/models/deepseek_v4/deepseek_v4_visual.py index 791361f911..3e730c5b2a 100644 --- a/lightllm/models/deepseek_v4/deepseek_v4_visual.py +++ b/lightllm/models/deepseek_v4/deepseek_v4_visual.py @@ -20,6 +20,7 @@ ) from lightllm.server.embed_cache.utils import get_shm_name_data, read_shm from lightllm.server.multimodal_params import ImageItem +from lightllm.server.visualserver import get_vit_attn_backend @lru_cache(8) @@ -72,18 +73,14 @@ def __init__(self, args): self.wqkv = nn.Linear(args.vision_dim, 3 * args.vision_dim, dtype=torch.bfloat16) self.wo = nn.Linear(args.vision_dim, args.vision_dim, dtype=torch.bfloat16) - def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, cu_seqlens: torch.Tensor) -> torch.Tensor: n = x.size(0) q, k, v = (t.view(n, self.n_heads, self.head_dim) for t in self.wqkv(x).chunk(3, dim=-1)) q = apply_rotary(q, cos, sin) k = apply_rotary(k, cos, sin) - # Keep the batch dimension so SDPA can select fused attention kernels. - o = F.scaled_dot_product_attention( - q.transpose(0, 1).unsqueeze(0), - k.transpose(0, 1).unsqueeze(0), - v.transpose(0, 1).unsqueeze(0), - ) - return self.wo(o.squeeze(0).transpose(0, 1).reshape(n, -1)) + o = torch.empty_like(q) + get_vit_attn_backend()(q, k, v, o, cu_seqlens, n) + return self.wo(o.reshape(n, -1)) class MLP(nn.Module): @@ -115,8 +112,8 @@ def __init__(self, args): self.norm2 = RMSNorm(args.vision_dim) self.mlp = MLP(args) - def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - x = x + self.attn(self.norm1(x), cos, sin) + def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, cu_seqlens: torch.Tensor) -> torch.Tensor: + x = x + self.attn(self.norm1(x), cos, sin, cu_seqlens) return x + self.mlp(self.norm2(x)) @@ -136,8 +133,9 @@ def forward(self, patches: torch.Tensor, n_h: int, n_w: int) -> torch.Tensor: cos, sin = get_vision_cos_sin(n_h, n_w, self.rope_dim, self.rope_theta) cos = cos.to(device=x.device) sin = sin.to(device=x.device) + cu_seqlens = torch.tensor([0, x.shape[0]], dtype=torch.int32, device=x.device) for block in self.blocks: - x = block(x, cos, sin) + x = block(x, cos, sin, cu_seqlens) return self.norm(x) From 32b57976439fc83fbff5b3a313678590619a173b Mon Sep 17 00:00:00 2001 From: shihaobai <42648726+shihaobai@users.noreply.github.com> Date: Mon, 7 Sep 2026 19:33:19 +0800 Subject: [PATCH 172/214] fix(pd): reserve transfer buffer during memory profiling (#1543) --- .../deepseek2_mem_manager.py | 12 ++---- .../deepseek4_mem_manager.py | 14 ++++++- .../kv_cache_mem_manager/mem_manager.py | 40 ++++++++++++++++--- 3 files changed, 51 insertions(+), 15 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py index 78e39cba2e..8cf2ac3301 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py @@ -28,14 +28,10 @@ def get_cell_size(self): def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.kv_buffer = torch.empty((layer_num, size + 1, head_num, head_dim), dtype=dtype, device="cuda") - def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - self.kv_move_buffer = torch.empty( - (page_num, page_size, self.layer_num, self.head_num, self.head_dim), dtype=self.dtype, device="cuda" - ) - self._buffer_mem_indexes_tensors = [ - torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num) - ] - return self.kv_move_buffer + def get_paged_kv_move_buffer_shape(self, page_num, page_size): + # DeepSeek MLA's single compressed KV latent is replicated across TP ranks, + # so the transfer buffer only needs one copy and does not multiply head_num by TP size. + return (page_num, page_size, self.layer_num, self.head_num, self.head_dim) def write_mem_to_page_kv_move_buffer( self, diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index a4c795a007..78d92fe251 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -7,7 +7,7 @@ from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator from lightllm.utils.dist_utils import get_current_rank_in_node -from lightllm.utils.envs_utils import get_unique_server_name +from lightllm.utils.envs_utils import get_env_start_args, get_unique_server_name from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -502,6 +502,18 @@ def get_fixed_memory_size(self): state_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 return self.n_c128 * state_rows * (2 * self.head_dim) * torch._utils._element_size(torch.float32) + def get_pd_kv_move_buffer_size(self): + args = get_env_start_args() + if args.run_mode not in ["prefill", "decode"]: + return 0 + layout = DeepseekV4PDCacheLayout.from_compress_rates( + self.compress_rates, + token_page_size=args.pd_kv_page_size, + head_dim=self.head_dim, + indexer_head_dim=self.indexer_head_dim, + ) + return args.pd_kv_page_num * layout.page_nbytes + # ------------------------------------------------------------------ buffers def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): rank_in_node = get_current_rank_in_node() diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index b2f13ecc04..7dcd71d127 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -1,5 +1,6 @@ -import re +import math import os +import re import torch import torch.distributed as dist import torch.multiprocessing as mp @@ -8,9 +9,12 @@ from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from .allocator import KvCacheAllocator from lightllm.utils.profile_max_tokens import get_available_gpu_memory, get_total_gpu_memory -from lightllm.utils.dist_utils import get_current_rank_in_node, get_node_world_size +from lightllm.utils.dist_utils import ( + get_current_device_id, + get_current_rank_in_node, + get_node_world_size, +) from lightllm.utils.envs_utils import get_unique_server_name, get_env_start_args -from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.config_utils import get_num_key_value_heads from lightllm.common.kv_trans_kernel.nixl_kv_trans import page_io from lightllm.utils.device_utils import kv_trans_use_p2p @@ -78,6 +82,26 @@ def get_cell_size(self): def get_fixed_memory_size(self): return 0 + def get_paged_kv_move_buffer_shape(self, page_num, page_size): + num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir) + return ( + page_num, + page_size, + self.layer_num, + 2 * num_kv_head, + self.head_dim, + ) + + def get_pd_kv_move_buffer_size(self): + args = get_env_start_args() + if args.run_mode not in ["prefill", "decode"]: + return 0 + shape = self.get_paged_kv_move_buffer_shape( + page_num=args.pd_kv_page_num, + page_size=args.pd_kv_page_size, + ) + return math.prod(shape) * torch._utils._element_size(self.dtype) + def profile_size(self, mem_fraction): if self.size is not None: return @@ -89,11 +113,15 @@ def profile_size(self, mem_fraction): fixed_memory_size = self.get_fixed_memory_size() reservations = getattr(self, "memory_reservations", {}) reserved_memory_size = sum(reservations.values()) - available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size + pd_kv_move_buffer_size = self.get_pd_kv_move_buffer_size() + available_memory_bytes = ( + available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size - pd_kv_move_buffer_size + ) if available_memory_bytes <= 0: raise RuntimeError( f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " f"plus {reserved_memory_size / 1024**3:.2f} GB reservations, " + f"plus {pd_kv_move_buffer_size / 1024**3:.2f} GB for the PD KV transfer buffer, " f"but only {available_memory:.2f} GB is available" ) self.size = int(available_memory_bytes / cell_size) @@ -105,6 +133,7 @@ def profile_size(self, mem_fraction): f"{str(available_memory)} GB space is available after load the model weight\n" f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" f"{reservations} bytes are reserved for post-profile model buffers\n" + f"{str(pd_kv_move_buffer_size / 1024 ** 2)} MB is reserved for PD KV transfer buffer\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" ) @@ -118,9 +147,8 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.kv_buffer = torch.empty((layer_num, size + 1, 2 * head_num, head_dim), dtype=dtype, device="cuda") def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor: - num_kv_head = get_num_key_value_heads(get_env_start_args().model_dir) self.kv_move_buffer = torch.empty( - (page_num, page_size, self.layer_num, 2 * num_kv_head, self.head_dim), dtype=self.dtype, device="cuda" + self.get_paged_kv_move_buffer_shape(page_num, page_size), dtype=self.dtype, device="cuda" ) self._buffer_mem_indexes_tensors = [ torch.empty((page_size,), dtype=torch.int64, device="cpu", pin_memory=True) for _ in range(page_num) From 02ad09efdd89f93efeb9cc53e5fe61a9c6a9c0da Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 11 Sep 2026 04:28:27 +0000 Subject: [PATCH 173/214] fix(dsv4): preallocate FlashMLA and C4 prefill workspaces Reuse caller-owned buffers to avoid repeated large CUDA allocations and long-running memory fragmentation. --- .../attention/nsa/dsv4_fp8_flashmla_sparse.py | 42 ++++- .../layer_infer/transformer_layer_infer.py | 133 ++++++++++++---- lightllm/models/deepseek_v4/model.py | 20 ++- .../triton_kernel/gather_c4_indexer_k_dsv4.py | 8 +- lightllm/models/deepseek_v4/workspace.py | 149 +++++++++++++++++- 5 files changed, 315 insertions(+), 37 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py index 33761d4eff..2205439b55 100644 --- a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py @@ -28,6 +28,34 @@ def _view_cache(buffer: torch.Tensor, page_size: int) -> torch.Tensor: return buffer[:, :byte_num].view(buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) +def _flashmla_sparse_decode_with_workspace(kwargs: dict, sched_meta, o_accum: torch.Tensor, lse_accum: torch.Tensor): + try: + op = torch.ops._flashmla_C.sparse_decode_fwd_with_workspace + except AttributeError as exc: + raise RuntimeError("DeepSeek-V4 prefill requires the workspace-enabled vllm._flashmla_C extension") from exc + + out, lse, new_tile_scheduler_metadata, new_num_splits = op( + kwargs["q"], + kwargs["k_cache"], + kwargs["indices"], + kwargs["topk_length"], + kwargs["attn_sink"], + sched_meta.tile_scheduler_metadata, + sched_meta.num_splits, + kwargs["extra_k_cache"], + kwargs["extra_indices_in_kvcache"], + kwargs["extra_topk_length"], + kwargs["head_dim_v"], + kwargs["softmax_scale"], + kwargs["out"], + o_accum, + lse_accum, + ) + sched_meta.tile_scheduler_metadata = new_tile_scheduler_metadata + sched_meta.num_splits = new_num_splits + return out, lse + + class DeepseekV4FlashMlaFp8SparseAttBackend(BaseAttBackend): def __init__(self, model): super().__init__(model=model) @@ -43,6 +71,8 @@ def _flashmla_att( nsa_dict: dict, sched_meta, flashmla_out: torch.Tensor = None, + flashmla_o_accum: torch.Tensor = None, + flashmla_lse_accum: torch.Tensor = None, ) -> torch.Tensor: from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( DSV4_C128_PAGE_SIZE, @@ -82,7 +112,15 @@ def _flashmla_att( ) if flashmla_out is not None: kwargs["out"] = flashmla_out - full_out, _ = flashmla.flash_mla_with_kvcache(**kwargs) + if flashmla_o_accum is None: + full_out, _ = flashmla.flash_mla_with_kvcache(**kwargs) + else: + full_out, _ = _flashmla_sparse_decode_with_workspace( + kwargs, + sched_meta, + flashmla_o_accum, + flashmla_lse_accum, + ) return full_out[:, 0, : self.real_q_head_num, :] def create_att_prefill_state(self, infer_state: "InferStateInfo") -> "_PrefillAttState": @@ -136,6 +174,8 @@ def prefill_att( nsa_dict, self._get_sched_meta(nsa_dict["compress_ratio"]), flashmla_out=full_out, + flashmla_o_accum=self.infer_state.dsv4_workspace.flashmla_prefill_o_accum, + flashmla_lse_accum=self.infer_state.dsv4_workspace.flashmla_prefill_lse_accum, ) if needs_padding: out.copy_(att_out) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index d3bfb5f3d7..6ac1e6af02 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -18,9 +18,7 @@ import deep_gemm from lightllm.models.deepseek_v4.triton_kernel.topk_transform import topk_transform_512 from lightllm.models.deepseek_v4.triton_kernel.topk_softplus_sqrt import topk_softplus_sqrt - - -_C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 +from lightllm.models.deepseek_v4.workspace import C4_LOGITS_ALIGNMENT, C4_PREFILL_LOGITS_BUDGET_BYTES class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer): @@ -432,10 +430,21 @@ def _compress_and_index(self, q_lora, x, infer_state: DeepseekV4InferStateInfo, x.record_stream(aux_stream) q_lora.record_stream(aux_stream) self.index_infer.write_indexer_k( - x, infer_state, layer_weight, cos_table, sin_table, use_custom_tensor_manager=False + x, + infer_state, + layer_weight, + cos_table, + sin_table, + use_custom_tensor_manager=False, + c4_aux_workspace=infer_state.dsv4_workspace.c4_prefill_aux, ) meta = self.index_infer.build_metadata( - x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=False + x, + q_lora, + infer_state, + layer_weight, + use_custom_tensor_manager=False, + c4_aux_workspace=infer_state.dsv4_workspace.c4_prefill_aux, ) self.compressor.compress(x, infer_state, layer_weight, cos_table, sin_table) main_stream.wait_stream(aux_stream) # join before prefill_att reads the indices / latent KV @@ -662,12 +671,17 @@ def compress( cos_table: torch.Tensor, sin_table: torch.Tensor, use_custom_tensor_manager: bool = True, + kv_score_out: torch.Tensor = None, + out_buffer: torch.Tensor = None, ): if self.compress_ratio == 0: return None if self.is_in_indexer: kv_score = layer_weight.idx_cmp_wkv_gate_.mm( - x, use_custom_tensor_mananger=use_custom_tensor_manager, out_dtype=torch.float32 + x, + out=kv_score_out, + use_custom_tensor_mananger=use_custom_tensor_manager, + out_dtype=torch.float32, ) norm_weight = layer_weight.idx_cmp_norm_.weight ape = layer_weight.idx_cmp_ape_.weight @@ -685,8 +699,7 @@ def compress( ape=ape, compress_ratio=self.compress_ratio, ) - out_buffer = None - if self.is_in_indexer and use_custom_tensor_manager: + if self.is_in_indexer and use_custom_tensor_manager and out_buffer is None: out_buffer = self.alloc_tensor( (infer_state.mem_index.numel(), self.index_head_dim), torch.bfloat16, @@ -747,10 +760,14 @@ def write_indexer_k( cos_table, sin_table, use_custom_tensor_manager=True, + c4_aux_workspace=None, ): if self.compress_ratio != 4: return # Only group-end rows in this dense bf16 scratch are valid indexer keys. + kv_score_out = indexer_k_out = hadamard_out = None + if c4_aux_workspace is not None: + kv_score_out, indexer_k_out, hadamard_out = c4_aux_workspace.indexer_k_buffers(x.shape[0]) scratch = self.indexer_compressor.compress( x, infer_state, @@ -758,12 +775,13 @@ def write_indexer_k( cos_table, sin_table, use_custom_tensor_manager=use_custom_tensor_manager, + kv_score_out=kv_score_out, + out_buffer=indexer_k_out, ) # Rotate K (post norm+rope) by the SAME 1/sqrt(d) Hadamard the q kernel applies, so # (Hq)·(Hk)=q·k (H orthogonal) and the fp8 quant of K stays accurate. from lightllm.models.deepseek3_2.triton_kernel.hadamard_transform import hadamard_transform - hadamard_out = None if use_custom_tensor_manager: hadamard_out = self.alloc_tensor(scratch.shape, scratch.dtype, device=scratch.device) scratch = hadamard_transform(scratch, scale=self.index_head_dim ** -0.5, out=hadamard_out) @@ -776,7 +794,13 @@ def write_indexer_k( ) def build_metadata( - self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + self, + x, + q_lora, + infer_state: DeepseekV4InferStateInfo, + layer_weight, + use_custom_tensor_manager=True, + c4_aux_workspace=None, ): swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1) swa_lengths = infer_state.dsv4_swa_lengths @@ -784,9 +808,16 @@ def build_metadata( extra_indices = extra_lengths = None if self.compress_ratio == 4: idx_q_fp8, weights = self._indexer_q_weight( - x, q_lora, infer_state, layer_weight, use_custom_tensor_manager=use_custom_tensor_manager + x, + q_lora, + infer_state, + layer_weight, + use_custom_tensor_manager=use_custom_tensor_manager, + c4_aux_workspace=c4_aux_workspace, + ) + extra_indices, extra_lengths = self._c4_indices( + infer_state, idx_q_fp8, weights, positions, c4_aux_workspace=c4_aux_workspace ) - extra_indices, extra_lengths = self._c4_indices(infer_state, idx_q_fp8, weights, positions) elif self.compress_ratio == 128: extra_indices = infer_state.dsv4_c128_indices.unsqueeze(1) extra_lengths = infer_state.dsv4_c128_lengths @@ -798,7 +829,13 @@ def build_metadata( } def _indexer_q_weight( - self, x, q_lora, infer_state: DeepseekV4InferStateInfo, layer_weight, use_custom_tensor_manager=True + self, + x, + q_lora, + infer_state: DeepseekV4InferStateInfo, + layer_weight, + use_custom_tensor_manager=True, + c4_aux_workspace=None, ): # Fused: wq_b mm -> rope(last rope dims) -> 1/sqrt(d) Hadamard -> per-token fp8 quant, with the # per-token q scale + indexer_weight_scale folded into weights, all in ONE kernel (was 4 kernels: @@ -813,14 +850,21 @@ def _indexer_q_weight( raise RuntimeError( f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}" ) - idx_q = layer_weight.idx_wq_b_.mm(q_lora, use_custom_tensor_mananger=use_custom_tensor_manager).view( - token_num, self.index_n_heads, self.index_head_dim - ) - raw_w = layer_weight.idx_weights_proj_.mm(x, use_custom_tensor_mananger=use_custom_tensor_manager).view( + idx_q_out = raw_w_out = None + if c4_aux_workspace is not None: + idx_q_out, raw_w_out = c4_aux_workspace.indexer_q_inputs(token_num) + idx_q = layer_weight.idx_wq_b_.mm( + q_lora, out=idx_q_out, use_custom_tensor_mananger=use_custom_tensor_manager + ).view(token_num, self.index_n_heads, self.index_head_dim) + raw_w = layer_weight.idx_weights_proj_.mm( + x, out=raw_w_out, use_custom_tensor_mananger=use_custom_tensor_manager + ).view( token_num, self.index_n_heads ) # [T, H] raw idx_q_fp8_out = weights_out = None - if use_custom_tensor_manager: + if c4_aux_workspace is not None: + idx_q_fp8_out, weights_out = c4_aux_workspace.indexer_q_outputs(token_num) + elif use_custom_tensor_manager: idx_q_fp8_out = self.alloc_tensor(idx_q.shape, torch.float8_e4m3fn, device=idx_q.device) weights_out = self.alloc_tensor((*idx_q.shape[:-1], 1), torch.float32, device=idx_q.device) idx_q_fp8, weights = fused_q_indexer_rope_hadamard_quant( @@ -834,7 +878,7 @@ def _indexer_q_weight( ) # fp8 [T,H,d]; weights [T,H,1] with q-scale + weight_scale folded return idx_q_fp8, weights.squeeze(-1) - def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions): + def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, positions, c4_aux_workspace=None): """c4 scorer via the page-safe deep_gemm.fp8_paged_mqa_logits over the paged c4 indexer pool, then masked topk-512 -> c4 slots. Fixed shapes (c4_cap pinned per graph bucket) keep the decode cuda graph capturable.""" @@ -859,7 +903,6 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, ) return slots.unsqueeze(1), lengths - device = positions.device page_size = mem_manager.c4_indexer_pool.page_size cached = getattr(infer_state, "_c4_paged_meta", None) @@ -868,7 +911,15 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, b_req_idx = infer_state.b_req_idx batch = b_req_idx.shape[0] - c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) # entries/req + token_num = positions.numel() + if c4_aux_workspace is None: + c4_len = torch.div(infer_state.b_seq_len, 4, rounding_mode="floor").to(torch.int32) + page_table_out = row_page_table = None + else: + page_cap = c4_cap // page_size + page_table_out, row_page_table = c4_aux_workspace.page_tables(batch, token_num, page_cap) + c4_len, _, _, _ = c4_aux_workspace.metadata(batch, token_num) + c4_len.copy_(torch.div(infer_state.b_seq_len, 4, rounding_mode="floor")) page_table = build_c4_indexer_page_table( mem_manager, b_req_idx, @@ -876,24 +927,35 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, c4_cap, infer_state.req_manager.req_to_token_indexs, infer_state.req_manager.HOLD_REQUEST_ID, + out=page_table_out, ) if infer_state.is_prefill: - token_batch_pos = torch.repeat_interleave( - torch.arange(batch, device=device, dtype=torch.int32), - infer_state.b_q_seq_len, - output_size=positions.numel(), - ) - row_page_table = page_table[token_batch_pos] + token_batch_pos = infer_state._dsv4_token_to_batch_idx + if row_page_table is None: + row_page_table = page_table[token_batch_pos] + else: + torch.index_select(page_table, 0, token_batch_pos, out=row_page_table) else: row_page_table = page_table - valid_len = ((positions + 1) // 4).to(torch.int32) - ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1) + if c4_aux_workspace is None: + valid_len = ((positions + 1) // 4).to(torch.int32) + ctx_lens = torch.clamp(valid_len, min=1).reshape(-1, 1) + topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) + else: + _, valid_len, ctx_lens, topk_lengths = c4_aux_workspace.metadata(batch, token_num) + valid_len.copy_((positions + 1) // 4) + torch.clamp(valid_len, min=1, out=ctx_lens[:, 0]) + torch.clamp(valid_len, min=1, max=index_topk, out=topk_lengths) rows_per_chunk = None chunk_metadata = None if infer_state.is_prefill: - rows_per_chunk = max(1, _C4_PREFILL_LOGITS_BUDGET_BYTES // (c4_cap * 4)) + if c4_aux_workspace is None: + aligned_c4_cap = ((c4_cap + C4_LOGITS_ALIGNMENT - 1) // C4_LOGITS_ALIGNMENT) * C4_LOGITS_ALIGNMENT + rows_per_chunk = max(1, C4_PREFILL_LOGITS_BUDGET_BYTES // (aligned_c4_cap * 4)) + else: + rows_per_chunk = c4_aux_workspace.rows_per_logits_chunk(c4_cap) if positions.numel() > rows_per_chunk: chunk_metadata = tuple( deep_gemm.get_paged_mqa_logits_metadata( @@ -913,7 +975,6 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, if chunk_metadata is None else None ) - topk_lengths = torch.clamp(torch.minimum(valid_len, torch.full_like(valid_len, index_topk)), min=1) cached = ( row_page_table, valid_len, @@ -949,6 +1010,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, valid_len[start:end], top_slots[start:end], page_size, + logits_out=(c4_aux_workspace.logits(end - start, c4_cap) if c4_aux_workspace is not None else None), ) return top_slots.unsqueeze(1), topk_lengths @@ -963,6 +1025,7 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, valid_len, top_slots, page_size, + logits_out=(c4_aux_workspace.logits(idx_q_fp8.shape[0], c4_cap) if c4_aux_workspace is not None else None), ) return top_slots.unsqueeze(1), topk_lengths @@ -978,8 +1041,9 @@ def _c4_score_topk( valid_len, top_slots, page_size, + logits_out=None, ): - logits = deep_gemm.fp8_paged_mqa_logits( + args = ( idx_q_fp8.unsqueeze(1), indexer_k_cache, weights, @@ -989,6 +1053,11 @@ def _c4_score_topk( c4_cap, False, ) + logits = ( + deep_gemm.fp8_paged_mqa_logits(*args) + if logits_out is None + else deep_gemm.fp8_paged_mqa_logits(*args, out=logits_out) + ) topk_transform_512( logits, valid_len, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 31afb8ed1f..09ed7a60ef 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -183,6 +183,11 @@ def _init_att_backend(self): head_dim_v=self.config["head_dim"], dtype=self.data_type, ) + if self.run_mode != "decode": + self.dsv4_workspace.init_flashmla_prefill_split_kv_workspace( + q_head_num=padded_q_head_num, + head_dim_v=self.config["head_dim"], + ) for layer_infer, layer_weight in zip(self.layers_infer, self.trans_layers_weight): layer_infer.flashmla_q_head_num_ = padded_q_head_num if padded_q_head_num == real_q_head_num: @@ -197,9 +202,22 @@ def _init_att_backend(self): def _init_custom(self): self._init_to_get_rotary() - self.dsv4_workspace = DeepseekV4Workspace(self) + prefill_aux_stream = None if os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1" and not self.args.enable_prefill_microbatch_overlap: prefill_aux_stream = torch.cuda.Stream() + self.dsv4_workspace = DeepseekV4Workspace(self) + if self.dsv4_workspace.needs_c4_prefill_aux(self): + import deep_gemm + + if "out:" not in (deep_gemm.fp8_paged_mqa_logits.__doc__ or ""): + raise RuntimeError( + "C4 prefill overlap workspace requires a DeepGEMM build whose " + "fp8_paged_mqa_logits API accepts out=" + ) + assert prefill_aux_stream is not None + with torch.cuda.stream(prefill_aux_stream): + self.dsv4_workspace.init_c4_prefill_aux(self) + if prefill_aux_stream is not None: for layer in self.layers_infer: layer.dsv4_prefill_aux_stream = prefill_aux_stream dist_group_manager.new_deepep_group( diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py index b50eab6726..54d2958836 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -42,6 +42,7 @@ def build_c4_indexer_page_table( c4_cap: int, req_to_token_indexs: torch.Tensor, hold_req_id: int, + out: torch.Tensor = None, ): """Build the logical-c4-page -> physical-c4-page table expected by DeepGEMM paged logits. @@ -54,7 +55,12 @@ def build_c4_indexer_page_table( assert c4_cap % page_size == 0 batch = b_req_idx.shape[0] page_cap = c4_cap // page_size - page_table = torch.empty((batch, page_cap), dtype=torch.int32, device=b_req_idx.device) + if out is None: + page_table = torch.empty((batch, page_cap), dtype=torch.int32, device=b_req_idx.device) + else: + assert out.shape == (batch, page_cap) + assert out.dtype == torch.int32 and out.device == b_req_idx.device and out.is_contiguous() + page_table = out _build_c4_indexer_page_table_kernel[(page_cap, batch)]( b_req_idx, c4_len, diff --git a/lightllm/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py index 4c8225b000..5aab324065 100644 --- a/lightllm/models/deepseek_v4/workspace.py +++ b/lightllm/models/deepseek_v4/workspace.py @@ -1,7 +1,118 @@ +import math +import os + import torch + +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C4_PAGE_SIZE from lightllm.utils.envs_utils import get_env_start_args +C4_PREFILL_LOGITS_BUDGET_BYTES = 512 * 1024 * 1024 +C4_LOGITS_ALIGNMENT = 256 + + +def _compress_cap(max_kv_seq_len: int, ratio: int) -> int: + entries = max(1, int(max_kv_seq_len) // ratio) + return ((entries + 63) // 64) * 64 + + +class DeepseekV4C4PrefillWorkspace: + """Persistent scratch used exclusively by the C4 prefill auxiliary stream.""" + + def __init__( + self, + token_capacity: int, + max_seq_length: int, + max_request_num: int, + index_n_heads: int, + index_head_dim: int, + ): + self.token_capacity = int(token_capacity) + self.max_request_num = int(max_request_num) + self.index_n_heads = int(index_n_heads) + self.index_head_dim = int(index_head_dim) + self.max_page_cap = _compress_cap(max_seq_length, 4) // DSV4_C4_PAGE_SIZE + + # The same byte arena serves the sequential indexer-K, indexer-Q and logits + # phases. Those phases run on one CUDA stream and therefore never overlap. + self.scratch = torch.empty((C4_PREFILL_LOGITS_BUDGET_BYTES,), dtype=torch.uint8, device="cuda") + self.idx_q_fp8 = torch.empty( + (self.token_capacity, self.index_n_heads, self.index_head_dim), + dtype=torch.float8_e4m3fn, + device="cuda", + ) + self.weights = torch.empty((self.token_capacity, self.index_n_heads, 1), dtype=torch.float32, device="cuda") + self.page_table = torch.empty((self.max_request_num * self.max_page_cap,), dtype=torch.int32, device="cuda") + self.row_page_table = torch.empty((self.token_capacity * self.max_page_cap,), dtype=torch.int32, device="cuda") + self.c4_len = torch.empty((self.max_request_num,), dtype=torch.int32, device="cuda") + self.valid_len = torch.empty((self.token_capacity,), dtype=torch.int32, device="cuda") + self.ctx_lens = torch.empty((self.token_capacity, 1), dtype=torch.int32, device="cuda") + self.topk_lengths = torch.empty((self.token_capacity,), dtype=torch.int32, device="cuda") + + max_indexer_k_phase = self.token_capacity * self.index_head_dim * (4 * 4 + 2 + 2) + max_indexer_q_phase = self.token_capacity * self.index_n_heads * (self.index_head_dim * 2 + 2) + assert max(max_indexer_k_phase, max_indexer_q_phase) <= C4_PREFILL_LOGITS_BUDGET_BYTES + + @staticmethod + def _flat_view(buffer: torch.Tensor, shape) -> torch.Tensor: + size = math.prod(shape) + assert size <= buffer.numel() + return buffer[:size].view(shape) + + def _scratch_view(self, shape, dtype: torch.dtype, byte_offset: int = 0) -> torch.Tensor: + nbytes = math.prod(shape) * torch._utils._element_size(dtype) + assert byte_offset % torch._utils._element_size(dtype) == 0 + assert byte_offset + nbytes <= self.scratch.numel() + return self.scratch[byte_offset : byte_offset + nbytes].view(dtype).view(shape) + + def indexer_k_buffers(self, token_num: int): + kv_score_shape = (token_num, 4 * self.index_head_dim) + indexer_k_shape = (token_num, self.index_head_dim) + kv_score = self._scratch_view(kv_score_shape, torch.float32) + offset = kv_score.numel() * kv_score.element_size() + indexer_k = self._scratch_view(indexer_k_shape, torch.bfloat16, offset) + offset += indexer_k.numel() * indexer_k.element_size() + hadamard = self._scratch_view(indexer_k_shape, torch.bfloat16, offset) + return kv_score, indexer_k, hadamard + + def indexer_q_inputs(self, token_num: int): + idx_q_shape = (token_num, self.index_n_heads * self.index_head_dim) + raw_weights_shape = (token_num, self.index_n_heads) + idx_q = self._scratch_view(idx_q_shape, torch.bfloat16) + offset = idx_q.numel() * idx_q.element_size() + raw_weights = self._scratch_view(raw_weights_shape, torch.bfloat16, offset) + return idx_q, raw_weights + + def indexer_q_outputs(self, token_num: int): + return self.idx_q_fp8[:token_num], self.weights[:token_num] + + def page_tables(self, batch_size: int, token_num: int, page_cap: int): + assert page_cap <= self.max_page_cap + page_table = self._flat_view(self.page_table, (batch_size, page_cap)) + row_page_table = self._flat_view(self.row_page_table, (token_num, page_cap)) + return page_table, row_page_table + + def metadata(self, batch_size: int, token_num: int): + return ( + self.c4_len[:batch_size], + self.valid_len[:token_num], + self.ctx_lens[:token_num], + self.topk_lengths[:token_num], + ) + + @staticmethod + def aligned_c4_cap(c4_cap: int) -> int: + return ((int(c4_cap) + C4_LOGITS_ALIGNMENT - 1) // C4_LOGITS_ALIGNMENT) * C4_LOGITS_ALIGNMENT + + def rows_per_logits_chunk(self, c4_cap: int) -> int: + return max(1, C4_PREFILL_LOGITS_BUDGET_BYTES // (self.aligned_c4_cap(c4_cap) * 4)) + + def logits(self, row_num: int, c4_cap: int) -> torch.Tensor: + aligned_c4_cap = self.aligned_c4_cap(c4_cap) + logits = self._scratch_view((row_num, aligned_c4_cap), torch.float32) + return logits[:, :c4_cap] + + class DeepseekV4Workspace: def __init__(self, model): self.token_capacity = int(model.batch_max_tokens) @@ -33,6 +144,20 @@ def __init__(self, model): self.c128_lengths = torch.empty((self.microbatch_count, self.token_capacity), dtype=torch.int32, device="cuda") self.flashmla_prefill_q = None self.flashmla_prefill_full_out = None + self.flashmla_prefill_o_accum = None + self.flashmla_prefill_lse_accum = None + self.c4_prefill_aux = None + + def init_c4_prefill_aux(self, model): + assert self.needs_c4_prefill_aux(model) + if self.c4_prefill_aux is None: + self.c4_prefill_aux = DeepseekV4C4PrefillWorkspace( + token_capacity=self.token_capacity, + max_seq_length=model.max_seq_length, + max_request_num=model.max_req_num, + index_n_heads=model.config["index_n_heads"], + index_head_dim=model.config["index_head_dim"], + ) def init_flashmla_prefill_q(self, real_q_head_num: int, padded_q_head_num: int, head_dim: int, dtype: torch.dtype): if self.flashmla_prefill_q is None: @@ -47,10 +172,30 @@ def init_flashmla_prefill_full_out(self, q_head_num: int, head_dim_v: int, dtype (self.token_capacity, 1, q_head_num, head_dim_v), dtype=dtype, device="cuda" ) + def init_flashmla_prefill_split_kv_workspace(self, q_head_num: int, head_dim_v: int): + if self.flashmla_prefill_o_accum is None: + sm_count = torch.cuda.get_device_properties().multi_processor_count + # FlashMLA stores one row per query plus at most one split row per SM. + split_capacity = self.token_capacity + sm_count + self.flashmla_prefill_o_accum = torch.empty( + (split_capacity, 1, q_head_num, head_dim_v), dtype=torch.float32, device="cuda" + ) + self.flashmla_prefill_lse_accum = torch.empty( + (split_capacity, 1, q_head_num), dtype=torch.float32, device="cuda" + ) + @staticmethod def compress_cap(max_kv_seq_len: int, ratio: int) -> int: - entries = max(1, int(max_kv_seq_len) // ratio) - return ((entries + 63) // 64) * 64 + return _compress_cap(max_kv_seq_len, ratio) + + @staticmethod + def needs_c4_prefill_aux(model) -> bool: + return ( + model.run_mode == "prefill" + and os.getenv("LIGHTLLM_DSV4_PREFILL_OVERLAP", "1") == "1" + and not model.args.enable_prefill_microbatch_overlap + and 4 in model.config["compress_ratios"] + ) def _alloc(self, width: int) -> torch.Tensor: return torch.empty((self.microbatch_count, self.token_capacity * width), dtype=torch.int32, device="cuda") From 62907acff6f10f08cb4112fbb7505cc2bc1a678f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 11 Sep 2026 04:43:06 +0000 Subject: [PATCH 174/214] add DISK_CACHE_NAME --- lightllm/server/api_start.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 06e15ed6f4..77c1e853ae 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -429,7 +429,10 @@ def _launch_subprocesses(args: StartArgs): instance_disk_cache_dir = None if args.enable_cpu_cache and args.enable_disk_cache: cache_base_dir = args.disk_cache_dir or tempfile.gettempdir() - instance_disk_cache_dir = os.path.join(cache_base_dir, f"lightllm_disk_cache_{get_unique_server_name()}") + disk_cache_name = os.getenv("DISK_CACHE_NAME") or f"lightllm_disk_cache_{get_unique_server_name()}" + if disk_cache_name in (".", "..") or os.path.basename(disk_cache_name) != disk_cache_name: + raise ValueError("DISK_CACHE_NAME must be a single directory name") + instance_disk_cache_dir = os.path.join(cache_base_dir, disk_cache_name) process_manager.register_disk_cache_dir(instance_disk_cache_dir) if args.enable_cpu_cache: From f14dd40738daf817d66ea8d025a93fe154fa8a8b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 14 Sep 2026 04:13:55 +0000 Subject: [PATCH 175/214] Structured output is not supported --- lightllm/server/core/objs/sampling_params.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index a136196141..b874631a91 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -380,11 +380,14 @@ def init(self, tokenizer, **kwargs): # Initialize guided_grammar guided_grammar = kwargs.get("guided_grammar", "") + guided_json = kwargs.get("guided_json", "") + if (guided_grammar or guided_json) and get_env_start_args().output_constraint_mode != "xgrammar": + raise ValueError("Structured output is not supported") + self.guided_grammar = GuidedGrammar() self.guided_grammar.initialize(guided_grammar, tokenizer) # Initialize guided_json - guided_json = kwargs.get("guided_json", "") self.guided_json = GuidedJsonSchema() self.guided_json.initialize(guided_json, tokenizer) From 2a7e6bd0f4ec6d2edda947c43c4d2d032ea97e09 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 14 Sep 2026 11:14:45 +0000 Subject: [PATCH 176/214] PD nodes unavailable: 503 --- lightllm/server/httpserver_for_pd_master/manager.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 53872127e0..0eabf7ec30 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -986,4 +986,10 @@ def update_node_load_info(self, load_info: Optional[dict]): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams ) -> Tuple[PD_Client_Obj, PD_Client_Obj, PDSelectionExtraInfo]: + if not self.prefill_nodes or not self.decode_nodes: + raise ServerBusyError( + "PD nodes unavailable: " + f"registered_prefill={len(self.prefill_nodes)}, registered_decode={len(self.decode_nodes)}", + status_code=503, + ) return self.selector.select_p_d_node(prompt, sampling_params, multimodal_params) From 33da96f290e6d58aba60feb658c8aecd84d61384 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 14 Sep 2026 12:58:58 +0000 Subject: [PATCH 177/214] fix(pd): harden control sends and reconnect cleanup --- lightllm/server/httpserver/pd_loop.py | 33 +++++++++++----- .../httpserver_for_pd_master/manager.py | 15 +++++--- lightllm/server/pd_io_struct.py | 38 +++++++++++++++++++ 3 files changed, 71 insertions(+), 15 deletions(-) diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index fb3c2a2b1f..db02f6b1fb 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -30,6 +30,9 @@ logger = init_logger(__name__) +_PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS = 5 +_PD_RECONNECT_DELAY_SECONDS = 10 + async def timer_log(manager: HttpServerManager): while True: @@ -83,10 +86,9 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O pd_handle_loop 主要负责与 pd master 进行注册连接,然后接收pd master发来的请求,然后 将推理结果转发给 pd master进行处理。 """ - # 创建转发队列 - forwarding_queue = AsyncQueue() - while True: + # 转发队列属于当前连接,避免超时未退出的旧请求在重连后上报过期 token。 + forwarding_queue = AsyncQueue() forwarding_tokens_task = None heartbeat_task = None generation_tasks: Dict[int, asyncio.Task] = {} @@ -144,9 +146,11 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O ) generation_tasks[group_req_id] = generation_task - def remove_generation_task(task: asyncio.Task, request_id: int = group_req_id): - if generation_tasks.get(request_id) is task: - generation_tasks.pop(request_id, None) + def remove_generation_task( + task: asyncio.Task, request_id: int = group_req_id, tasks=generation_tasks + ): + if tasks.get(request_id) is task: + tasks.pop(request_id, None) # task 可能在首次运行前被取消,此时协程内的 finally 不会执行。 manager.cancel_pd_request_registration(request_id) @@ -184,10 +188,21 @@ def remove_generation_task(task: asyncio.Task, request_id: int = group_req_id): for task in child_tasks: task.cancel() if child_tasks: - await asyncio.gather(*child_tasks, return_exceptions=True) + done_tasks, pending_tasks = await asyncio.wait( + child_tasks, timeout=_PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS + ) + if done_tasks: + await asyncio.gather(*done_tasks, return_exceptions=True) + if pending_tasks: + logger.warning( + "timed out after %s seconds cleaning up %s PD child task(s); reconnecting", + _PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS, + len(pending_tasks), + ) + for task in pending_tasks: + task.cancel() - await asyncio.sleep(10) - await forwarding_queue.get_all_data() + await asyncio.sleep(_PD_RECONNECT_DELAY_SECONDS) logger.info("reconnection to pd_master") diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 0eabf7ec30..9bdb758397 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -449,7 +449,7 @@ async def fetch_pd_stream( old_max_new_tokens = sampling_params.max_new_tokens sampling_params.max_new_tokens = 1 - await p_node.websocket.send_bytes(pickle.dumps((ObjType.REQ, (prompt, sampling_params, multimodal_params)))) + await p_node.send_control_message(pickle.dumps((ObjType.REQ, (prompt, sampling_params, multimodal_params)))) try: await self._wait_for_event_or_disconnect( @@ -468,7 +468,7 @@ async def fetch_pd_stream( logger.info(f"group_request_id: {group_request_id} get prefill prompt ids len {len(prompt_ids)}") sampling_params.max_new_tokens = old_max_new_tokens - await d_node.websocket.send_bytes( + await d_node.send_control_message( pickle.dumps((ObjType.REQ, (prompt_ids, sampling_params, MultimodalParams()))) ) @@ -489,7 +489,7 @@ async def fetch_pd_stream( upkv_status: PDUpKVStatus = up_status_event.upkv_status pd_kv_trans_params: bytes = upkv_status.pd_kv_trans_params decode_node_info: PDDecodeNodeInfo = pickle.loads(pd_kv_trans_params) - await p_node.websocket.send_bytes( + await p_node.send_control_message( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) @@ -676,12 +676,12 @@ async def abort( pass try: - await p_node.websocket.send_bytes(pickle.dumps((ObjType.ABORT, group_request_id))) + await p_node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) except: pass try: - await d_node.websocket.send_bytes(pickle.dumps((ObjType.ABORT, group_request_id))) + await d_node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) except: pass @@ -955,7 +955,10 @@ def register_pd(self, pd_info_json, websocket): def remove_pd(self, pd_info_json): pd_client = PD_Client_Obj(**pd_info_json) - self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) + removed_client = self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) + if removed_client is not None: + # In-flight requests can still hold this node after it leaves the selector. + removed_client.websocket = None self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index ddf1487e96..b5093c1be8 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -1,3 +1,4 @@ +import asyncio import enum import time import copy @@ -161,6 +162,8 @@ class PD_Client_Obj: dispatched_prompt_chars: int = 0 # 当前派发到该节点且尚未产出首 token 的请求数。 dispatched_req_num: int = 0 + _send_lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False, repr=False, compare=False) + _send_task: Optional[asyncio.Task] = field(default=None, init=False, repr=False, compare=False) def __post_init__(self): if self.mode not in ["prefill", "decode"]: @@ -172,6 +175,41 @@ def __post_init__(self): def to_llm_url(self): return f"http://{self.client_ip_port}/pd_generate_stream" + async def send_control_message(self, payload: bytes) -> None: + # Waiting requests remain cancellable BEFORE they advance the compression dictionary. + await self._send_lock.acquire() + try: + if self.websocket is None: + raise ConnectionError(f"PD control connection unavailable: {self.client_ip_port}") + send_task = asyncio.create_task(self.websocket.send_bytes(payload)) + self._send_task = send_task + except BaseException: + self._send_lock.release() + raise + + def finish_send(task: asyncio.Task): + self._send_task = None + try: + task.result() + except BaseException: + self.websocket = None + logger.exception("PD control send failed: peer=%s", self.client_ip_port) + finally: + self._send_lock.release() + + # The connection owns the task AND the lock until the complete frame is sent. + # Hypercorn compresses before awaiting its TCP send lock. Cancelling that wait + # drops the frame but leaves the deflate dictionary advanced for later messages. + send_task.add_done_callback(finish_send) + try: + await asyncio.shield(send_task) + except asyncio.CancelledError: + logger.warning( + "PD control send caller cancelled; connection-owned send continues: " "peer=%s", + self.client_ip_port, + ) + raise + @dataclass class PD_Master_Obj: From 9a3f99dcf3c492797105362d229479e29db71211 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 15 Sep 2026 01:56:43 +0000 Subject: [PATCH 178/214] fix(pd): clean up requests correctly after node disconnect - stop requests using the disconnected node - keep newly reconnected nodes registered - process decode tasks before their abort messages --- lightllm/server/api_http_pd.py | 4 +-- lightllm/server/httpserver/pd_loop.py | 4 +++ .../httpserver_for_pd_master/manager.py | 27 ++++++++++--------- lightllm/server/pd_io_struct.py | 5 ++++ .../decode_node_impl/decode_trans_process.py | 11 +++++--- 5 files changed, 33 insertions(+), 18 deletions(-) diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 54ff2ab5b5..500ac3df3e 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -33,7 +33,7 @@ async def register_and_keep_alive(websocket: WebSocket): logger.info(f"Client connected from IP: {client_ip}, Port: {client_port}") regist_json = json.loads(await websocket.receive_text()) logger.info(f"received regist_json {regist_json}") - await g_objs.httpserver_manager.register_pd(regist_json, websocket) + pd_client = await g_objs.httpserver_manager.register_pd(regist_json, websocket) try: heartbeat_timeout_seconds = 30 @@ -60,7 +60,7 @@ async def register_and_keep_alive(websocket: WebSocket): logger.exception(str(e)) finally: logger.error(f"client {regist_json} removed") - await g_objs.httpserver_manager.remove_pd(regist_json) + await g_objs.httpserver_manager.remove_pd(pd_client) return diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index db02f6b1fb..450552a62a 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -183,6 +183,10 @@ def remove_generation_task( logger.error("connetion to pd_master has error") logger.exception(str(e)) finally: + # Cancel the connection's requests even if their generators cannot exit promptly. + # abort() also defers cancellation for requests that have not registered shm_req yet. + for group_req_id in generation_tasks: + await manager.abort(group_req_id) child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] child_tasks.extend(generation_tasks.values()) for task in child_tasks: diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 9bdb758397..582caa138a 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -89,11 +89,14 @@ def is_healthy(self): return False async def register_pd(self, pd_info_json, websocket): - self.pd_manager.register_pd(pd_info_json, websocket) - return - - async def remove_pd(self, pd_info_json): - self.pd_manager.remove_pd(pd_info_json) + return self.pd_manager.register_pd(pd_info_json, websocket) + + async def remove_pd(self, pd_client: PD_Client_Obj): + self.pd_manager.remove_pd(pd_client) + # Wake every stage so the request's existing error path aborts the surviving peer. + for req_status in self.req_id_to_out_inf.values(): + if req_status.p_node is pd_client or req_status.d_node is pd_client: + await req_status.set_error(f"PD {pd_client.mode} node {pd_client.client_ip_port} disconnected") return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -950,15 +953,15 @@ def register_pd(self, pd_info_json, websocket): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} registed") - return + return pd_client - def remove_pd(self, pd_info_json): - pd_client = PD_Client_Obj(**pd_info_json) + def remove_pd(self, pd_client: PD_Client_Obj): + # A closing connection must not remove a newer registration at the same address. + pd_client.websocket = None + if self.url_to_pd_nodes.get(pd_client.client_ip_port) is not pd_client: + return - removed_client = self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) - if removed_client is not None: - # In-flight requests can still hold this node after it leaves the selector. - removed_client.websocket = None + self.url_to_pd_nodes.pop(pd_client.client_ip_port) self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index b5093c1be8..e713b1bae2 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -176,6 +176,11 @@ def to_llm_url(self): return f"http://{self.client_ip_port}/pd_generate_stream" async def send_control_message(self, payload: bytes) -> None: + # A disconnected client may still have an old send holding the lock. Do not + # let cleanup messages wait for that send before noticing the invalidation. + if self.websocket is None: + raise ConnectionError(f"PD control connection unavailable: {self.client_ip_port}") + # Waiting requests remain cancellable BEFORE they advance the compression dictionary. await self._send_lock.acquire() try: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index c1ec78835c..21a792fd39 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -193,10 +193,10 @@ def _warmup(self): def recv_task_loop(self): while True: obj: Union[PDChunckedTransTaskGroup, PDAbortReq] = self.task_in_queue.get() - if isinstance(obj, PDChunckedTransTaskGroup): + if isinstance(obj, (PDChunckedTransTaskGroup, PDAbortReq)): + # Keep the producer's group-before-abort order through dispatch as well. + # Aborting here can miss a group that is still waiting in this queue. self.recv_task_group_queue.put(obj) - elif isinstance(obj, PDAbortReq): - self._abort(request_id=obj.request_id) else: assert False, f"recv error obj {obj}" @@ -217,7 +217,10 @@ def _abort(self, request_id: int, error_info: str = "aborted req"): @log_exception def dispatch_task_loop(self): while True: - trans_task_group: PDChunckedTransTaskGroup = self.recv_task_group_queue.get() + trans_task_group: Union[PDChunckedTransTaskGroup, PDAbortReq] = self.recv_task_group_queue.get() + if isinstance(trans_task_group, PDAbortReq): + self._abort(request_id=trans_task_group.request_id) + continue with self.waiting_dict_lock: for task in trans_task_group.task_list: From 54c57304dc105be7e9cea6379c40ad682cf5e68d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 15 Sep 2026 02:51:39 +0000 Subject: [PATCH 179/214] unit test --- .../httpserver/test_pd_generate_error.py | 51 +++++++++++++++---- unit_tests/server/test_pd_master_mode.py | 17 +++++++ 2 files changed, 57 insertions(+), 11 deletions(-) diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 3c15fe8766..438bcf6d84 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -12,7 +12,7 @@ HttpServerManagerForPDMaster, ReqStatus, ) -from lightllm.server.pd_io_struct import ObjType +from lightllm.server.pd_io_struct import ObjType, PD_Client_Obj from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError @@ -253,8 +253,8 @@ async def run(): manager.metric_client = MagicMock() manager.infos_queues = None - p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) - d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + p_node = SimpleNamespace(send_control_message=AsyncMock()) + d_node = SimpleNamespace(send_control_message=AsyncMock()) req_status = ReqStatus(123, p_node, d_node) manager.req_id_to_out_inf = {123: req_status} @@ -269,8 +269,8 @@ async def run(): assert req_status.prefill_prompt_ids_event.is_set() assert req_status.up_status_event.is_set() assert manager.req_id_to_out_inf[123] is req_status - p_node.websocket.send_bytes.assert_not_awaited() - d_node.websocket.send_bytes.assert_not_awaited() + p_node.send_control_message.assert_not_awaited() + d_node.send_control_message.assert_not_awaited() with pytest.raises( RuntimeError, @@ -353,8 +353,8 @@ def test_pd_master_abort_removes_request_even_when_node_notifications_fail(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.req_id_to_out_inf = {} - p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock(side_effect=ConnectionError("p down")))) - d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock(side_effect=ConnectionError("d down")))) + p_node = SimpleNamespace(send_control_message=AsyncMock(side_effect=ConnectionError("p down"))) + d_node = SimpleNamespace(send_control_message=AsyncMock(side_effect=ConnectionError("d down"))) manager.req_id_to_out_inf[123] = ReqStatus(123, p_node, d_node) await manager.abort(123) @@ -368,12 +368,41 @@ def test_pd_master_abort_uses_explicit_nodes_when_request_status_is_missing(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.req_id_to_out_inf = {} - p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) - d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + p_node = SimpleNamespace(send_control_message=AsyncMock()) + d_node = SimpleNamespace(send_control_message=AsyncMock()) await manager.abort(123, p_node=p_node, d_node=d_node) - p_node.websocket.send_bytes.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) - d_node.websocket.send_bytes.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) + p_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) + d_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) + + asyncio.run(run()) + + +def test_pd_master_abort_does_not_wait_for_disconnected_node_inflight_send(): + async def run(): + send_started = asyncio.Event() + release_send = asyncio.Event() + + class _BlockingWebSocket: + async def send_bytes(self, _payload): + send_started.set() + await release_send.wait() + + p_node = PD_Client_Obj(1, "prefill:8000", "prefill", {}, websocket=_BlockingWebSocket()) + d_node = SimpleNamespace(send_control_message=AsyncMock()) + inflight_send = asyncio.create_task(p_node.send_control_message(b"request")) + await send_started.wait() + + manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) + manager.req_id_to_out_inf = {123: ReqStatus(123, p_node, d_node)} + p_node.websocket = None + + try: + await asyncio.wait_for(manager.abort(123), timeout=1) + d_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) + finally: + release_send.set() + await inflight_send asyncio.run(run()) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 1c704f6bc8..c044aedc6f 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -8,6 +8,7 @@ from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager +from lightllm.utils.error_utils import ServerBusyError def test_pd_node_self_request_limit_cli_defaults_to_disabled_and_can_be_enabled(): @@ -207,6 +208,22 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True +@pytest.mark.parametrize( + ("prefill_nodes", "decode_nodes"), + [([], [object()]), ([object()], [])], +) +def test_pd_manager_returns_service_unavailable_when_a_node_role_is_empty(prefill_nodes, decode_nodes): + manager = PDManager(StartArgs()) + manager.prefill_nodes = prefill_nodes + manager.decode_nodes = decode_nodes + + with pytest.raises(ServerBusyError) as exc_info: + manager.select_p_d_node("prompt", None, None) + + assert exc_info.value.status_code == 503 + assert "PD nodes unavailable" in exc_info.value.message + + def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) From 2f37b957a6d9f83394dddc13a3e94e5dbc8b3b1a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 15 Sep 2026 10:52:49 +0000 Subject: [PATCH 180/214] add convert_deepseek_v4_scales_to_ue8m0_fp32 --- ...onvert_deepseek_v4_scales_to_ue8m0_fp32.py | 337 ++++++++++++++++++ 1 file changed, 337 insertions(+) create mode 100755 tools/convert_deepseek_v4_scales_to_ue8m0_fp32.py diff --git a/tools/convert_deepseek_v4_scales_to_ue8m0_fp32.py b/tools/convert_deepseek_v4_scales_to_ue8m0_fp32.py new file mode 100755 index 0000000000..0a5cb7b8c5 --- /dev/null +++ b/tools/convert_deepseek_v4_scales_to_ue8m0_fp32.py @@ -0,0 +1,337 @@ +#!/usr/bin/env python3 +"""Convert DeepSeek-V4 block-FP8 scales to UE8M0-equivalent FP32 values. + +The output scale tensors remain FP32, but every value is rounded upward to a +power of two with the same rule used by LightLLM's ``USE_UE8M0_SCALE`` path:: + + scale = ceil_to_power_of_two(max(amax, 1e-4) / 448) + +Only scales paired with 2-D F8_E4M3 weights and matching the block-128 layout +are converted. Paired FP8 weights and all unrelated tensors are preserved. +The source model is never modified. + +Rounding a non-power-of-two scale without requantizing its paired FP8 weight +changes the effective dequantized weight. This is intentional for this tool. +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import shutil +from dataclasses import dataclass +from pathlib import Path +from typing import Dict, Mapping, Sequence, Tuple + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + + +FP8_MAX = 448.0 +MIN_AMAX = 1.0e-4 + + +@dataclass(frozen=True) +class ScaleSpec: + name: str + source_shard: str + source_dtype: str + shape: Tuple[int, ...] + + +def _load_index(model_dir: Path) -> Tuple[Path, dict]: + index_path = model_dir / "model.safetensors.index.json" + if not index_path.is_file(): + raise FileNotFoundError(f"missing safetensors index: {index_path}") + with index_path.open("r", encoding="utf-8") as file: + index = json.load(file) + if not isinstance(index.get("weight_map"), dict): + raise ValueError(f"invalid weight_map in {index_path}") + return index_path, index + + +def _inspect_tensor_metadata(model_dir: Path, weight_map: Mapping[str, str]) -> Dict[str, Tuple[str, Tuple[int, ...]]]: + names_by_shard: Dict[str, list[str]] = {} + for name, shard in weight_map.items(): + names_by_shard.setdefault(shard, []).append(name) + + metadata: Dict[str, Tuple[str, Tuple[int, ...]]] = {} + for shard, names in names_by_shard.items(): + shard_path = model_dir / shard + if not shard_path.is_file(): + raise FileNotFoundError(f"missing shard: {shard_path}") + with safe_open(shard_path, framework="pt", device="cpu") as file: + shard_keys = set(file.keys()) + missing = set(names) - shard_keys + if missing: + raise ValueError(f"{shard} is missing indexed tensors: {sorted(missing)[:5]}") + for name in names: + tensor_slice = file.get_slice(name) + metadata[name] = ( + tensor_slice.get_dtype(), + tuple(int(dim) for dim in tensor_slice.get_shape()), + ) + return metadata + + +def _scale_candidates(weight_name: str) -> Tuple[str, str]: + base = weight_name[: -len(".weight")] + return base + ".scale", base + ".weight_scale_inv" + + +def _inspect_scale_specs( + weight_map: Mapping[str, str], + tensor_metadata: Mapping[str, Tuple[str, Tuple[int, ...]]], + block_size: int, +) -> list[ScaleSpec]: + specs: list[ScaleSpec] = [] + for weight_name, (weight_dtype, weight_shape) in tensor_metadata.items(): + if weight_dtype != "F8_E4M3" or not weight_name.endswith(".weight"): + continue + if len(weight_shape) != 2: + raise ValueError(f"FP8 weight must be 2-D: {weight_name} has shape {weight_shape}") + + scale_names = [name for name in _scale_candidates(weight_name) if name in tensor_metadata] + if len(scale_names) != 1: + raise ValueError(f"expected one paired scale for {weight_name}, found {scale_names}") + scale_name = scale_names[0] + scale_dtype, scale_shape = tensor_metadata[scale_name] + if scale_dtype not in {"F32", "F8_E8M0"}: + raise ValueError(f"{scale_name} has unsupported dtype {scale_dtype}; expected F32 or F8_E8M0") + + expected_shape = ( + math.ceil(weight_shape[0] / block_size), + math.ceil(weight_shape[1] / block_size), + ) + if scale_shape != expected_shape: + raise ValueError(f"{scale_name} shape is {scale_shape}, expected {expected_shape}") + specs.append( + ScaleSpec( + name=scale_name, + source_shard=weight_map[scale_name], + source_dtype=scale_dtype, + shape=scale_shape, + ) + ) + + if not specs: + raise ValueError("no block-FP8 scale tensors found") + return specs + + +def ceil_to_ue8m0_fp32(scale: torch.Tensor) -> torch.Tensor: + """Return FP32 powers of two using LightLLM's UE8M0 rounding rule.""" + + scale = scale.to(dtype=torch.float32, device="cpu").contiguous() + if not bool(torch.isfinite(scale).all()): + raise ValueError("scale contains a non-finite value") + if bool((scale < 0).any()): + raise ValueError("scale contains a negative value") + + scale = scale.clamp_min(MIN_AMAX / FP8_MAX) + bits = scale.view(torch.int32) + exponent = ((bits >> 23) & 0xFF) + ((bits & 0x7FFFFF) != 0).to(torch.int32) + exponent = exponent.clamp_(1, 254) + return (exponent << 23).view(torch.float32) + + +def _find_specs_needing_conversion(model_dir: Path, specs: Sequence[ScaleSpec]) -> list[ScaleSpec]: + specs_by_shard: Dict[str, list[ScaleSpec]] = {} + for spec in specs: + specs_by_shard.setdefault(spec.source_shard, []).append(spec) + + needed: list[ScaleSpec] = [] + for shard, shard_specs in specs_by_shard.items(): + with safe_open(model_dir / shard, framework="pt", device="cpu") as file: + for spec in shard_specs: + source_scale = file.get_tensor(spec.name) + converted_scale = ceil_to_ue8m0_fp32(source_scale) + if source_scale.dtype != torch.float32 or not torch.equal(source_scale, converted_scale): + needed.append(spec) + return needed + + +def _copy_or_link(source: Path, destination: Path, mode: str) -> None: + if mode == "copy": + shutil.copy2(source, destination) + return + if mode == "symlink": + destination.symlink_to(source.resolve()) + return + try: + os.link(source, destination) + except OSError: + shutil.copy2(source, destination) + + +def _prepare_output_dir( + source_dir: Path, + output_dir: Path, + *, + affected_shards: set[str], + link_mode: str, +) -> None: + if output_dir.exists(): + if any(output_dir.iterdir()): + raise FileExistsError(f"output directory is not empty: {output_dir}") + else: + output_dir.mkdir(parents=True) + + for entry in source_dir.iterdir(): + if not entry.is_file() or entry.name == "model.safetensors.index.json": + continue + if entry.suffix == ".safetensors" and entry.name in affected_shards: + continue + _copy_or_link(entry, output_dir / entry.name, link_mode) + + +def _rewrite_shard( + source_path: Path, + output_path: Path, + scale_names: set[str], +) -> Tuple[int, int]: + changed_tensors = 0 + changed_values = 0 + temporary_path = output_path.with_name(f".{output_path.name}.{os.getpid()}.tmp") + try: + with safe_open(source_path, framework="pt", device="cpu") as source_file: + source_keys = set(source_file.keys()) + missing = scale_names - source_keys + if missing: + raise ValueError(f"{source_path.name} is missing scales: {sorted(missing)[:5]}") + + tensors = {} + for name in source_file.keys(): + tensor = source_file.get_tensor(name) + if name in scale_names: + converted = ceil_to_ue8m0_fp32(tensor) + changed = int(torch.count_nonzero(converted != tensor.to(torch.float32)).item()) + changed_values += changed + changed_tensors += int(changed > 0 or tensor.dtype != torch.float32) + tensor = converted + tensors[name] = tensor.contiguous() + metadata = source_file.metadata() + save_file(tensors, temporary_path, metadata=metadata) + shutil.copymode(source_path, temporary_path) + os.replace(temporary_path, output_path) + finally: + temporary_path.unlink(missing_ok=True) + return changed_tensors, changed_values + + +def _verify_output(output_dir: Path, specs: Sequence[ScaleSpec]) -> None: + specs_by_shard: Dict[str, list[ScaleSpec]] = {} + for spec in specs: + specs_by_shard.setdefault(spec.source_shard, []).append(spec) + + for shard, shard_specs in specs_by_shard.items(): + with safe_open(output_dir / shard, framework="pt", device="cpu") as file: + for spec in shard_specs: + tensor_slice = file.get_slice(spec.name) + if tensor_slice.get_dtype() != "F32": + raise ValueError(f"{spec.name} is {tensor_slice.get_dtype()}, expected F32") + if tuple(tensor_slice.get_shape()) != spec.shape: + raise ValueError(f"{spec.name} shape changed") + scale = file.get_tensor(spec.name) + if not bool(torch.isfinite(scale).all()) or not bool((scale > 0).all()): + raise ValueError(f"{spec.name} contains an invalid scale") + log2_scale = torch.log2(scale) + if not torch.equal(log2_scale, torch.round(log2_scale)): + raise ValueError(f"{spec.name} contains a non-power-of-two value") + + +def convert(args: argparse.Namespace) -> None: + source_dir = Path(args.source_model_dir).expanduser().resolve(strict=True) + output_dir = Path(args.output_model_dir).expanduser().resolve(strict=False) + if source_dir == output_dir: + raise ValueError("source and output directories must be different") + if args.block_size <= 0: + raise ValueError("--block-size must be positive") + + _, index = _load_index(source_dir) + weight_map: Dict[str, str] = index["weight_map"] + tensor_metadata = _inspect_tensor_metadata(source_dir, weight_map) + specs = _inspect_scale_specs(weight_map, tensor_metadata, args.block_size) + needed_specs = _find_specs_needing_conversion(source_dir, specs) + specs_by_shard: Dict[str, set[str]] = {} + for spec in needed_specs: + specs_by_shard.setdefault(spec.source_shard, set()).add(spec.name) + + source_dtype_counts: Dict[str, int] = {} + for spec in needed_specs: + source_dtype_counts[spec.source_dtype] = source_dtype_counts.get(spec.source_dtype, 0) + 1 + rewrite_bytes = sum((source_dir / shard).stat().st_size for shard in specs_by_shard) + print(f"source: {source_dir}") + print(f"output: {output_dir}") + print(f"block-FP8 scales: {len(specs)}; requiring conversion: {len(needed_specs)}") + print(f"source dtypes requiring conversion: {source_dtype_counts}") + print(f"affected shards: {len(specs_by_shard)}; rewrite size: {rewrite_bytes / 1024**3:.2f} GiB") + print("output scales: FP32 powers of two; paired FP8 weights are not requantized") + if args.dry_run: + return + + _prepare_output_dir( + source_dir, + output_dir, + affected_shards=set(specs_by_shard), + link_mode=args.link_mode, + ) + + changed_tensors = 0 + changed_values = 0 + for index_in_plan, shard in enumerate(sorted(specs_by_shard), start=1): + print(f"[{index_in_plan}/{len(specs_by_shard)}] rewriting {shard}", flush=True) + shard_changed_tensors, shard_changed_values = _rewrite_shard( + source_dir / shard, + output_dir / shard, + specs_by_shard[shard], + ) + changed_tensors += shard_changed_tensors + changed_values += shard_changed_values + + metadata = dict(index.get("metadata") or {}) + if "total_size" in metadata: + added_bytes = sum(math.prod(spec.shape) * 3 for spec in needed_specs if spec.source_dtype == "F8_E8M0") + metadata["total_size"] = int(metadata["total_size"]) + added_bytes + output_index = {**index, "metadata": metadata, "weight_map": weight_map} + index_path = output_dir / "model.safetensors.index.json" + temporary_index_path = output_dir / ".model.safetensors.index.json.tmp" + with temporary_index_path.open("w", encoding="utf-8") as file: + json.dump(output_index, file, ensure_ascii=False, indent=2) + file.write("\n") + os.replace(temporary_index_path, index_path) + + if not args.no_verify: + print("verifying converted scales ...", flush=True) + _verify_output(output_dir, specs) + print( + f"conversion complete: {changed_tensors} scale tensors and " + f"{changed_values} values changed; output: {output_dir}" + ) + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("source_model_dir", help="source Hugging Face model directory") + parser.add_argument("output_model_dir", help="new output model directory") + parser.add_argument("--block-size", type=int, default=128, help="FP8 block size (default: 128)") + parser.add_argument( + "--link-mode", + choices=("hardlink", "copy", "symlink"), + default="hardlink", + help="how unchanged model files and shards are placed in the output directory", + ) + parser.add_argument("--dry-run", action="store_true", help="inspect and print the conversion plan only") + parser.add_argument("--no-verify", action="store_true", help="skip final scale verification") + return parser + + +def main() -> None: + convert(build_parser().parse_args()) + + +if __name__ == "__main__": + main() From f5f3ed2cc74857ef8821840d5ca0d42c4f2a3e67 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 15 Sep 2026 13:04:36 +0000 Subject: [PATCH 181/214] align B300 --- lightllm/common/quantization/deepgemm.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 436d0bc50f..1217a2bc74 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -5,7 +5,6 @@ from lightllm.common.quantization.registry import QUANTMETHODS from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import per_token_group_quant_fp8 from lightllm.utils.log_utils import init_logger -from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -71,7 +70,7 @@ def quantize(self, weight: torch.Tensor, output: WeightPack): from lightllm.common.basemodel.triton_kernel.quantization.fp8w8a8_block_quant_kernel import weight_quant device = output.weight.device - weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=is_sm100_gpu()) + weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=True) output.weight.copy_(weight) output.weight_scale.copy_(scale) return @@ -99,7 +98,7 @@ def apply( column_major_scales=True, scale_tma_aligned=True, alloc_func=alloc_func, - use_ue8m0_scales=is_sm100_gpu(), + use_ue8m0_scales=True, ) if out is None: From 7cc102740b41a7c985a0cac0022e952634d0a796 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 16 Sep 2026 00:41:15 +0000 Subject: [PATCH 182/214] fix has_nvlink --- lightllm/utils/device_utils.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/lightllm/utils/device_utils.py b/lightllm/utils/device_utils.py index 7135e3f232..8cedb001a7 100644 --- a/lightllm/utils/device_utils.py +++ b/lightllm/utils/device_utils.py @@ -150,8 +150,9 @@ def has_nvlink(): # Call nvidia-smi to get the topology matrix result = subprocess.check_output(["nvidia-smi", "topo", "--matrix"]) result = result.decode("utf-8") - # Check if the output contains 'NVLink' - return any(f"NV{i}" in result for i in range(1, 8)) + # NVLink topology entries are reported as NV followed by the link count, + # for example NV8 on H800 and NV18 on B300. + return any(entry.startswith("NV") and entry[2:].isdigit() for entry in result.split()) except FileNotFoundError: # nvidia-smi is not installed, assume no NVLink return False From c899b318b29786a39dbb9dd2bc63e186e3779687 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 16 Sep 2026 09:13:04 +0000 Subject: [PATCH 183/214] add triton ep backend --- ...8192,topk=6,world_size=8}_NVIDIA_H200.json | 16 + ...8192,topk=6,world_size=8}_NVIDIA_H200.json | 16 + .../fused_moe/fused_moe_weight.py | 4 +- .../meta_weights/fused_moe/impl/__init__.py | 5 +- .../fused_moe/impl/triton_ep_impl.py | 42 ++ .../fused_moe/grouped_fused_moe_ep.py | 5 +- .../fused_moe/sm90_fp8_triton_ep_moe.py | 575 ++++++++++++++++++ .../quantization/fp8act_quant_kernel.py | 36 +- lightllm/common/quantization/deepgemm.py | 4 +- lightllm/distributed/communication_op.py | 77 ++- .../layer_infer/transformer_layer_infer.py | 17 +- .../layer_infer/transformer_layer_infer.py | 13 +- .../layer_infer/transformer_layer_infer.py | 17 +- lightllm/server/api_cli.py | 11 + lightllm/server/api_start.py | 9 +- lightllm/server/core/objs/start_args_type.py | 1 + .../mode_backend/ep_balance_monitor.py | 1 + 17 files changed, 787 insertions(+), 62 deletions(-) create mode 100644 lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json create mode 100644 lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py create mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py diff --git a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json new file mode 100644 index 0000000000..694771f831 --- /dev/null +++ b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json @@ -0,0 +1,16 @@ +{ + "publish_rows_large": 32, + "publish_rows_small": 4, + "publish_num_warps": 2, + "histogram_block": 2048, + "histogram_num_warps": 8, + "pull_num_warps": 8, + "activation_programs": 1024, + "activation_rows": 32, + "activation_num_warps": 8, + "return_num_warps": 8, + "combine_rows_large": 2, + "combine_rows_small": 1, + "combine_block": 512, + "combine_num_warps": 4 +} diff --git a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json new file mode 100644 index 0000000000..625756625c --- /dev/null +++ b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json @@ -0,0 +1,16 @@ +{ + "publish_rows_large": 32, + "publish_rows_small": 4, + "publish_num_warps": 8, + "histogram_block": 512, + "histogram_num_warps": 2, + "pull_num_warps": 8, + "activation_programs": 1024, + "activation_rows": 32, + "activation_num_warps": 8, + "return_num_warps": 8, + "combine_rows_large": 2, + "combine_rows_small": 1, + "combine_block": 512, + "combine_num_warps": 4 +} diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 60bcb3a39d..cbfd109cca 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -61,7 +61,8 @@ def __init__( self.moe_intermediate_size = moe_intermediate_size self.quant_method = quant_method assert num_fused_shared_experts in [0, 1], "num_fused_shared_experts can only support 0 or 1 now." - self.enable_ep_moe = get_env_start_args().enable_ep_moe + args = get_env_start_args() + self.enable_ep_moe = args.enable_ep_moe self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self._init_config(network_config) @@ -73,6 +74,7 @@ def __init__( routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, expert_parallel_state=self.expert_parallel_state, + ep_moe_backend=args.ep_moe_backend, ) self.lock = threading.Lock() self._moe_weight_finalized = False diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 89fc529801..96e060faca 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -1,5 +1,6 @@ from lightllm.common.quantization.quantize_method import QuantizationMethod from .triton_impl import FuseMoeTriton +from .triton_ep_impl import FuseMoeTritonEP from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM from .mxfp4_impl import FuseMoeMXFP4 @@ -13,13 +14,15 @@ def create_fuse_moe_impl( routed_scaling_factor: float, quant_method: QuantizationMethod, expert_parallel_state: ExpertParallelState | None = None, + ep_moe_backend: str = "auto", ): if quant_method.method_name == "mxfp4w4a16-b32-marlin": if expert_parallel_state is not None: raise RuntimeError("mxfp4w4a16-b32-marlin does not support enable_ep_moe yet") impl_cls = FuseMoeMXFP4 elif expert_parallel_state is not None: - impl_cls = FuseMoeDeepGEMM + use_triton_ep = ep_moe_backend == "triton" and quant_method.method_name == "fp8w8a8-b128-deepgemm" + impl_cls = FuseMoeTritonEP if use_triton_ep else FuseMoeDeepGEMM elif quant_method.method_name == "awq_marlin": impl_cls = FuseMoeMarlin else: diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py new file mode 100644 index 0000000000..edfc95a8a0 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py @@ -0,0 +1,42 @@ +from typing import Optional + +import torch + +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( + ExpertParallelState, +) +from lightllm.common.quantization.quantize_method import WeightPack +from lightllm.distributed import dist_group_manager + +from .triton_impl import FuseMoeTriton + + +class FuseMoeTritonEP(FuseMoeTriton): + """Triton MoE backend for expert-parallel symmetric-memory execution.""" + + def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): + super().__init__(*args, **kwargs) + self.expert_parallel_state = expert_parallel_state + + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, + ) -> torch.Tensor: + buffer = dist_group_manager.ep_triton_moe_buffer + return buffer.forward( + input_tensor, + (w13.weight, w13.weight_scale), + (w2.weight, w2.weight_scale), + topk_weights, + topk_ids.to(torch.long), + float("inf") if clamp_limit is None else clamp_limit, + alloc_tensor_func=alloc_tensor_func, + ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index a8b28ceaa3..7495d69f72 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -130,7 +130,10 @@ def _get_mega_moe_cumulative_stats(num_local_experts: int, device: torch.device, return stats -def prepare_mega_moe_weights(w13: Any, w2: Any, quant_method: Any): +def prepare_ep_moe_weights(w13: Any, w2: Any, quant_method: Any): + if dist_group_manager.ep_triton_moe_quant_method == quant_method.method_name: + quant_method.mega_moe_mma_type = None + return mma_type = dist_group_manager.ep_mega_moe_mma_type if dist_group_manager.ep_mega_moe_quant_method != quant_method.method_name: quant_method.mega_moe_mma_type = None diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py b/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py new file mode 100644 index 0000000000..a380f0c02e --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py @@ -0,0 +1,575 @@ +"""Single-node SM90 FP8 expert parallelism over CUDA symmetric memory.""" + +from typing import Callable, Optional, Tuple + +import deep_gemm +import torch +import torch.distributed as dist +import torch.distributed._symmetric_memory as symm +import triton +import triton.language as tl +from frozendict import frozendict + +from lightllm.common.kernel_config import KernelConfigs + + +DEFAULT_CONFIG = { + "publish_rows_large": 16, + "publish_rows_small": 4, + "publish_num_warps": 4, + "histogram_block": 1024, + "histogram_num_warps": 4, + "pull_num_warps": 4, + "activation_programs": 512, + "activation_rows": 16, + "activation_num_warps": 4, + "return_num_warps": 4, + "combine_rows_large": 4, + "combine_rows_small": 1, + "combine_block": 512, + "combine_num_warps": 4, +} + + +class SM90FP8TritonEPMoEKernelConfig(KernelConfigs): + kernel_name = "sm90_fp8_triton_ep_moe" + + @classmethod + def _params( + cls, + hidden_size: int, + intermediate_size: int, + local_experts: int, + topk: int, + world_size: int, + num_max_tokens_per_rank: int, + alignment: int, + ): + return frozendict( + { + "hidden_size": hidden_size, + "intermediate_size": intermediate_size, + "local_experts": local_experts, + "topk": topk, + "world_size": world_size, + "num_max_tokens_per_rank": num_max_tokens_per_rank, + "alignment": alignment, + } + ) + + @classmethod + def try_to_get_best_config(cls, **kwargs) -> dict: + config = cls.get_the_config(cls._params(**kwargs)) + return dict(DEFAULT_CONFIG if config is None else config) + + @classmethod + def get_config_if_available(cls, **kwargs) -> Optional[dict]: + config = cls.get_the_config(cls._params(**kwargs)) + return None if config is None else dict(config) + + @classmethod + def save_config(cls, config: dict, **kwargs) -> None: + cls.store_config(cls._params(**kwargs), config) + + +@triton.jit +def _publish( + X, + IDS, + WEIGHTS, + Q, + SF, + SID, + SW, + ROWS, + COUNTS, + M, + H: tl.constexpr, + K: tl.constexpr, + E: tl.constexpr, + BE: tl.constexpr, + BK: tl.constexpr, + BM: tl.constexpr, +): + rows = tl.program_id(0) * BM + tl.arange(0, BM) + block = tl.program_id(1) + cols = block * 128 + tl.arange(0, 128) + values = tl.load(X + rows[:, None] * H + cols[None, :], rows[:, None] < M, 0).to(tl.float32) + scale = tl.maximum(tl.max(tl.abs(values), 1), 1e-10) / 448.0 + quant = tl.clamp(tl.div_rn(values, scale[:, None]), -448.0, 448.0).to(Q.dtype.element_ty) + tl.store(Q + rows[:, None] * H + cols[None, :], quant, rows[:, None] < M) + tl.store(SF + rows * (H // 128) + block, scale, rows < M) + if block == 0: + slots = tl.arange(0, BK) + offsets = rows[:, None] * K + slots[None, :] + mask = (rows[:, None] < M) & (slots[None, :] < K) + tl.store(SID + offsets, tl.load(IDS + offsets, mask, -1), mask) + tl.store(SW + offsets, tl.load(WEIGHTS + offsets, mask, 0), mask) + if tl.program_id(0) == 0: + tl.store(ROWS, M) + experts = tl.arange(0, BE) + tl.store(COUNTS + experts, 0, experts < E) + + +@triton.jit +def _histogram( + ID_PTRS, + ROW_PTRS, + COUNTS, + CAP: tl.constexpr, + K: tl.constexpr, + E: tl.constexpr, + RANK: tl.constexpr, + BE: tl.constexpr, + BLOCK: tl.constexpr, +): + peer = tl.program_id(1) + ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) + rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + expert = tl.load(ids + offsets, offsets < rows * K, -1).to(tl.int32) - RANK * E + valid = (offsets < rows * K) & (expert >= 0) & (expert < E) + counts = tl.histogram(expert, BE, mask=valid) + experts = tl.arange(0, BE) + tl.atomic_add(COUNTS + experts, counts, experts < E, sem="relaxed") + + +@triton.jit +def _prefix(COUNTS, ENDS, CURSOR, E: tl.constexpr, BE: tl.constexpr, ALIGN: tl.constexpr): + experts = tl.arange(0, BE) + counts = tl.load(COUNTS + experts, experts < E, 0) + padded = tl.cdiv(counts, ALIGN) * ALIGN + ends = tl.cumsum(padded) + tl.store(ENDS + experts, ends, experts < E) + tl.store(CURSOR + experts, ends - padded, experts < E) + + +@triton.jit +def _pull( + X_PTRS, + SF_PTRS, + ID_PTRS, + ROW_PTRS, + CURSOR, + X, + SF, + ROW_MAP, + CAP: tl.constexpr, + H: tl.constexpr, + K: tl.constexpr, + E: tl.constexpr, + RANK: tl.constexpr, + R: tl.constexpr, + WORLD: tl.constexpr, + BK: tl.constexpr, +): + token = tl.program_id(0) // WORLD + peer = (tl.program_id(0) % WORLD + RANK) % WORLD + rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) + if token < rows: + ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) + slots = tl.arange(0, BK) + experts = tl.load(ids + token * K + slots, slots < K, -1).to(tl.int32) - RANK * E + local = (slots < K) & (experts >= 0) & (experts < E) + if tl.sum(local.to(tl.int32), 0) > 0: + source_x = tl.load(X_PTRS + peer).to(tl.pointer_type(tl.float8e4nv)) + source_sf = tl.load(SF_PTRS + peer).to(tl.pointer_type(tl.float32)) + cols = tl.arange(0, H) + scale_cols = tl.arange(0, H // 128) + values = tl.load(source_x + token * H + cols) + scales = tl.load(source_sf + token * (H // 128) + scale_cols) + for slot in range(K): + expert = tl.load(ids + token * K + slot).to(tl.int32) - RANK * E + if (expert >= 0) & (expert < E): + row = tl.atomic_add(CURSOR + expert, 1, sem="relaxed") + tl.store(X + row * H + cols, values) + tl.store(SF + row + scale_cols * R, scales) + tl.store(ROW_MAP + (peer * CAP + token) * K + slot, row) + + +@triton.jit +def _activation( + X, + Q, + SF, + ENDS, + E: tl.constexpr, + R: tl.constexpr, + I: tl.constexpr, + LIMIT: tl.constexpr, + BM: tl.constexpr, +): + total = tl.load(ENDS + E - 1) + blocks_n = I // 128 + for tile in range(tl.program_id(0), tl.cdiv(total, BM) * blocks_n, tl.num_programs(0)): + row = (tile // blocks_n) * BM + tl.arange(0, BM) + block = tile % blocks_n + col = block * 128 + tl.arange(0, 128) + offset = row[:, None] * (2 * I) + col[None, :] + gate = tl.load(X + offset, row[:, None] < total, 0).to(tl.float32) + up = tl.load(X + offset + I, row[:, None] < total, 0) + gate = tl.minimum(gate, LIMIT) + up = tl.clamp(up, -LIMIT, LIMIT) + gate = (gate / (1 + tl.exp(-gate))).to(tl.bfloat16) + act = (up * gate).to(tl.bfloat16).to(tl.float32) + scale = tl.maximum(tl.max(tl.abs(act), 1), 1e-10) / 448.0 + quant = tl.clamp(act / scale[:, None], -448.0, 448.0).to(Q.dtype.element_ty) + tl.store(Q + row[:, None] * I + col[None, :], quant, row[:, None] < total) + tl.store(SF + row + block * R, scale, row < total) + + +@triton.jit +def _return( + X, + ROW_MAP, + ID_PTRS, + WEIGHT_PTRS, + ROW_PTRS, + OUT_PTRS, + CAP: tl.constexpr, + H: tl.constexpr, + K: tl.constexpr, + E: tl.constexpr, + RANK: tl.constexpr, + WORLD: tl.constexpr, +): + token = tl.program_id(0) // WORLD + peer = (tl.program_id(0) % WORLD + RANK) % WORLD + rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) + if token < rows: + ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) + weights = tl.load(WEIGHT_PTRS + peer).to(tl.pointer_type(tl.float32)) + cols = tl.arange(0, H) + acc = tl.full((H,), 0.0, tl.float32) + first = -1 + for slot in range(K): + expert = tl.load(ids + token * K + slot).to(tl.int32) - RANK * E + if (expert >= 0) & (expert < E): + row = tl.load(ROW_MAP + (peer * CAP + token) * K + slot) + weight = tl.load(weights + token * K + slot) + acc += tl.load(X + row * H + cols).to(tl.float32) * weight + if first < 0: + first = slot + if first >= 0: + output = tl.multiple_of(tl.load(OUT_PTRS + peer).to(tl.pointer_type(tl.bfloat16)), 16) + tl.store(output + (token * K + first) * H + cols, acc.to(tl.bfloat16)) + + +@triton.jit +def _combine( + RETURNED, + IDS, + Y, + M, + H: tl.constexpr, + K: tl.constexpr, + E: tl.constexpr, + BLOCK: tl.constexpr, + BM: tl.constexpr, +): + row = tl.program_id(0) * BM + tl.arange(0, BM) + col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) + acc = tl.full((BM, BLOCK), 0, tl.float32) + for slot in tl.static_range(K): + expert = tl.load(IDS + row * K + slot, row < M, -1) + valid = (row < M) & (expert >= 0) + for earlier in tl.static_range(slot): + previous = tl.load(IDS + row * K + earlier, row < M, -1) + valid &= (previous < 0) | (previous // E != expert // E) + values = tl.load( + RETURNED + (row[:, None] * K + slot) * H + col[None, :], + valid[:, None], + 0, + ).to(tl.float32) + acc += values + tl.store(Y + row[:, None] * H + col[None, :], acc.to(tl.bfloat16), row[:, None] < M) + + +class SM90FP8TritonEPMoEBuffer: + """SM90 FP8 Triton EP MoE 的共享通信与计算缓冲区。 + + 每个 EP group 只创建一个实例,并由所有 MoE 层复用;各层的专家权重仍由 layer + weight 持有。每个 rank 将本地 token、路由结果发布到 symmetric memory,本 rank + 拉取所有发往本地专家的 token,完成两次 grouped GEMM,再把加权后的局部结果写回 + token 所在 rank。当前生产选择器只在单节点 SM90 Prefill 路径启用该类。 + """ + + def __init__( + self, + group, + num_experts: int, + num_max_tokens_per_rank: int, + topk: int, + hidden_size: int, + intermediate_size: int, + alignment: Optional[int] = None, + ): + self.group = group + self.rank = dist.get_rank(group) + self.world = dist.get_world_size(group) + assert num_experts % self.world == 0 + self.num_experts = num_experts + self.local_experts = num_experts // self.world + self.num_max_tokens_per_rank = num_max_tokens_per_rank + self.topk = topk + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + + # 普通尾块使用 128 对齐;完整 num_max_tokens_per_rank 块若有单独调优配置, + # 可以切到 256 对齐。显式传入 alignment 时则始终使用同一套配置。 + self.alignment = 128 if alignment is None else alignment + assert hidden_size == 2 * intermediate_size + assert hidden_size % 128 == 0 and intermediate_size % 128 == 0 + assert self.alignment in (128, 256) + self.config = SM90FP8TritonEPMoEKernelConfig.try_to_get_best_config( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + local_experts=self.local_experts, + topk=topk, + world_size=self.world, + num_max_tokens_per_rank=num_max_tokens_per_rank, + alignment=self.alignment, + ) + self.full_alignment = self.alignment + self.full_config = self.config + if alignment is None: + full_config = SM90FP8TritonEPMoEKernelConfig.get_config_if_available( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + local_experts=self.local_experts, + topk=topk, + world_size=self.world, + num_max_tokens_per_rank=num_max_tokens_per_rank, + alignment=256, + ) + if full_config is not None: + self.full_alignment = 256 + self.full_config = full_config + + # 每个本地专家的接收段都需要独立补齐,最坏情况下额外占用 + # local_experts * (alignment - 1) 行。 + max_recv_rows = self.world * num_max_tokens_per_rank * topk + workspace_alignment = max(self.alignment, self.full_alignment) + self.workspace_rows = ( + triton.cdiv(max_recv_rows + self.local_experts * (workspace_alignment - 1), workspace_alignment) + * workspace_alignment + ) + self.handles = [] + self.pointers = [] + self.sources = [] + # sources/pointers 的下标是后续 Triton kernel 的固定协议: + # 0: FP8 hidden states,1: 每 128 列一组的量化 scale + # 2: top-k expert id,3: top-k weight,4: 本 rank 实际 token 数 + # 5: 各专家 owner 写回的 BF16 局部结果 + for shape, dtype in ( + ((num_max_tokens_per_rank, hidden_size), torch.float8_e4m3fn), + ((num_max_tokens_per_rank, hidden_size // 128), torch.float32), + ((num_max_tokens_per_rank, topk), torch.int64), + ((num_max_tokens_per_rank, topk), torch.float32), + ((1,), torch.int32), + ((num_max_tokens_per_rank, topk, hidden_size), torch.bfloat16), + ): + tensor = symm.empty(*shape, dtype=dtype, device="cuda") + handle = symm.rendezvous(tensor, group=group) + assert all(pointer % 16 == 0 for pointer in handle.buffer_ptrs) + self.sources.append(tensor) + self.handles.append(handle) + self.pointers.append(torch.tensor(handle.buffer_ptrs, device="cuda", dtype=torch.int64)) + + # counts 记录每个本地专家收到的 route 数;ends 是对齐后的排他结束位置; + # cursor 从各专家段起点开始原子递增;row_map 保存 (peer, token, top-k slot) + # 到本地 expert-major workspace 行号的映射,供结果回传使用。 + self.counts = torch.empty(self.local_experts, device="cuda", dtype=torch.int32) + self.ends = torch.empty_like(self.counts) + self.cursor = torch.empty_like(self.counts) + self.row_map = torch.empty(self.world * num_max_tokens_per_rank * topk, device="cuda", dtype=torch.int32) + + # x/x_scale 保存按本地专家分段后的 FP8 输入;workspace 依次承载 W1 和 W2 + # 的 BF16 输出。W1 完成后 x 已不再使用,且 H=2I,因此其前半空间可原地 + # 复用为量化后的激活输入,避免再分配一份 FP8 activation buffer。 + self.x = torch.empty(self.workspace_rows, hidden_size, device="cuda", dtype=torch.float8_e4m3fn) + self.x_scale = torch.empty(hidden_size // 128, self.workspace_rows, device="cuda", dtype=torch.float32).T + self.workspace = torch.empty(self.workspace_rows, hidden_size, device="cuda", dtype=torch.bfloat16) + self.activation = self.x.view(self.workspace_rows * 2, intermediate_size)[: self.workspace_rows] + self.activation_scale = self.x_scale[:, : intermediate_size // 128] + + def _get_runtime_alignment_and_config(self, rows: int) -> Tuple[int, dict]: + """完整块使用 full-chunk 调优结果,尾块使用更稳妥的基础配置。""" + if rows == self.num_max_tokens_per_rank: + return self.full_alignment, self.full_config + return self.alignment, self.config + + @property + def allocated_bytes(self) -> int: + """返回本 rank 主要通信与工作区 tensor 的字节数,供基准测试统计。""" + tensors = self.sources + [ + self.counts, + self.ends, + self.cursor, + self.row_map, + self.x, + self.x_scale, + self.workspace, + ] + return sum(tensor.numel() * tensor.element_size() for tensor in tensors) + + def forward( + self, + hidden_states: torch.Tensor, + w1: Tuple[torch.Tensor, torch.Tensor], + w2: Tuple[torch.Tensor, torch.Tensor], + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + clamp_limit: float, + alloc_tensor_func: Callable = torch.empty, + ) -> torch.Tensor: + """执行一次 EP MoE,并返回与 hidden_states 同 shape、同 dtype 的结果。""" + rows = hidden_states.shape[0] + assert rows <= self.num_max_tokens_per_rank + assert hidden_states.shape == (rows, self.hidden_size) + assert hidden_states.dtype == torch.bfloat16 and hidden_states.is_contiguous() + assert topk_ids.shape == topk_weights.shape == (rows, self.topk) + assert topk_ids.dtype == torch.int64 and topk_ids.is_contiguous() + assert topk_weights.dtype == torch.float32 and topk_weights.is_contiguous() + assert w1[0].shape == (self.local_experts, 2 * self.intermediate_size, self.hidden_size) + assert w2[0].shape == (self.local_experts, self.hidden_size, self.intermediate_size) + + alignment, config = self._get_runtime_alignment_and_config(rows) + publish_rows = config["publish_rows_large"] if rows >= 16 else config["publish_rows_small"] + + # 1. 量化本地输入并发布输入、scale 和路由元数据。barrier 之后所有 rank + # 才能安全读取彼此的 symmetric-memory source buffers。 + _publish[(max(1, triton.cdiv(rows, publish_rows)), self.hidden_size // 128)]( + hidden_states, + topk_ids, + topk_weights, + *self.sources[:5], + self.counts, + rows, + self.hidden_size, + self.topk, + self.local_experts, + triton.next_power_of_2(self.local_experts), + triton.next_power_of_2(self.topk), + publish_rows, + num_warps=config["publish_num_warps"], + ) + self.handles[0].barrier(channel=0, timeout_ms=10000) + + # 2. 统计所有 peer 发往本地专家的 route,生成对齐后的 expert-major 分段, + # 再从 peer buffer 拉取输入并记录每条 route 在 workspace 中的行号。 + histogram_block = config["histogram_block"] + _histogram[(triton.cdiv(self.num_max_tokens_per_rank * self.topk, histogram_block), self.world)]( + self.pointers[2], + self.pointers[4], + self.counts, + self.num_max_tokens_per_rank, + self.topk, + self.local_experts, + self.rank, + triton.next_power_of_2(self.local_experts), + histogram_block, + num_warps=config["histogram_num_warps"], + ) + _prefix[(1,)]( + self.counts, + self.ends, + self.cursor, + self.local_experts, + triton.next_power_of_2(self.local_experts), + alignment, + ) + _pull[(self.num_max_tokens_per_rank * self.world,)]( + self.pointers[0], + self.pointers[1], + self.pointers[2], + self.pointers[4], + self.cursor, + self.x, + self.x_scale, + self.row_map, + self.num_max_tokens_per_rank, + self.hidden_size, + self.topk, + self.local_experts, + self.rank, + self.workspace_rows, + self.world, + triton.next_power_of_2(self.topk), + num_warps=config["pull_num_warps"], + ) + + # 3. 在本地 expert-major 布局上执行 W1 -> SwiGLU+FP8 quant -> W2。 + # DeepGEMM 的 contiguous-layout alignment 是进程级状态,调用后必须恢复。 + expected_m = triton.cdiv(self.num_max_tokens_per_rank * self.topk, self.local_experts) + previous_alignment = deep_gemm.get_mk_alignment_for_contiguous_layout() + deep_gemm.set_mk_alignment_for_contiguous_layout(alignment) + try: + deep_gemm.m_grouped_fp8_gemm_nt_contiguous( + (self.x, self.x_scale), + w1, + self.workspace, + self.ends, + use_psum_layout=True, + expected_m_for_psum_layout=expected_m, + ) + _activation[(config["activation_programs"],)]( + self.workspace, + self.activation, + self.activation_scale, + self.ends, + self.local_experts, + self.workspace_rows, + self.intermediate_size, + clamp_limit, + config["activation_rows"], + num_warps=config["activation_num_warps"], + ) + deep_gemm.m_grouped_fp8_gemm_nt_contiguous( + (self.activation, self.activation_scale), + w2, + self.workspace, + self.ends, + use_psum_layout=True, + expected_m_for_psum_layout=expected_m, + ) + finally: + deep_gemm.set_mk_alignment_for_contiguous_layout(previous_alignment) + + # 4. 每个专家 owner 按 row_map 找回源 token,在本地先合并并乘 route weight, + # 然后向源 rank 写回一份局部结果。第二次 barrier 保证写回全部可见。 + _return[(self.num_max_tokens_per_rank * self.world,)]( + self.workspace, + self.row_map, + self.pointers[2], + self.pointers[3], + self.pointers[4], + self.pointers[5], + self.num_max_tokens_per_rank, + self.hidden_size, + self.topk, + self.local_experts, + self.rank, + self.world, + num_warps=config["return_num_warps"], + ) + self.handles[5].barrier(channel=0, timeout_ms=10000) + + # 5. 一个 token 可能命中多个 owner rank;源 rank 对各 owner 的局部结果求和。 + output = alloc_tensor_func(hidden_states.shape, device=hidden_states.device, dtype=hidden_states.dtype) + if rows: + combine_rows = config["combine_rows_large"] if rows >= 16 else config["combine_rows_small"] + combine_block = config["combine_block"] + _combine[(triton.cdiv(rows, combine_rows), self.hidden_size // combine_block)]( + self.sources[5], + self.sources[2], + output, + rows, + self.hidden_size, + self.topk, + self.local_experts, + combine_block, + combine_rows, + num_warps=config["combine_num_warps"], + ) + return output diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py index c3a5b9f634..553d80e586 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py @@ -29,6 +29,7 @@ def _per_token_group_quant_fp8( y_q_ptr, y_s_ptr, y_stride, + M, N, eps, fp8_min, @@ -48,18 +49,19 @@ def _per_token_group_quant_fp8( NEED_MASK: tl.constexpr, USE_UE8M0_SCALE: tl.constexpr, COPY_TOPK: tl.constexpr, + GROUPS_PER_CTA: tl.constexpr, ): - g_id = tl.program_id(0) - y_ptr += g_id * y_stride - y_q_ptr += g_id * y_stride + g_id = tl.program_id(0) * GROUPS_PER_CTA + tl.arange(0, GROUPS_PER_CTA) + y_ptr += g_id[:, None] * y_stride + y_q_ptr += g_id[:, None] * y_stride row_id = g_id // xs_n col_id = g_id % xs_n y_s_ptr += row_id * xs_stride_m + col_id * xs_stride_n - cols = tl.arange(0, BLOCK) # N <= BLOCK + cols = tl.arange(0, BLOCK)[None, :] # N <= BLOCK - if NEED_MASK: - mask = cols < N + if NEED_MASK or GROUPS_PER_CTA > 1: + mask = (g_id[:, None] < M) & (cols < N) other = 0.0 else: mask = None @@ -67,21 +69,21 @@ def _per_token_group_quant_fp8( y = tl.load(y_ptr + cols, mask=mask, other=other).to(tl.float32) # Quant - _absmax = tl.max(tl.abs(y)) + _absmax = tl.max(tl.abs(y), axis=1) if USE_UE8M0_SCALE: y_s = _ceil_to_ue8m0(tl.maximum(_absmax, 1.0e-4) / fp8_max) else: y_s = tl.maximum(_absmax, eps) / fp8_max - y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty) + y_q = tl.clamp(y / y_s[:, None], fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty) tl.store(y_q_ptr + cols, y_q, mask=mask) - tl.store(y_s_ptr, y_s) + tl.store(y_s_ptr, y_s, mask=g_id < M) if COPY_TOPK: - topk_cols = tl.arange(0, TOPK_BLOCK) - topk_mask = (col_id == 0) & (topk_cols < num_topk) - topk_offsets = row_id * topk_row_stride + topk_cols - topk_out_offsets = row_id * topk_out_row_stride + topk_cols + topk_cols = tl.arange(0, TOPK_BLOCK)[None, :] + topk_mask = (g_id[:, None] < M) & (col_id[:, None] == 0) & (topk_cols < num_topk) + topk_offsets = row_id[:, None] * topk_row_stride + topk_cols + topk_out_offsets = row_id[:, None] * topk_out_row_stride + topk_cols topk_ids = tl.load(topk_ids_ptr + topk_offsets, mask=topk_mask) topk_weights = tl.load(topk_weights_ptr + topk_offsets, mask=topk_mask) tl.store(topk_ids_out_ptr + topk_out_offsets, topk_ids, mask=topk_mask) @@ -136,11 +138,16 @@ def lightllm_per_token_group_quant_fp8( topk_ids = topk_weights = topk_ids_out = topk_weights_out = x topk_block = 1 num_topk = topk_row_stride = topk_out_row_stride = 0 - _per_token_group_quant_fp8[(M,)]( + # Large Mega inputs otherwise launch one CTA per 128-element group. + groups_per_cta = 16 if copy_topk and M >= 4096 else 1 + if groups_per_cta > 1: + num_warps = 4 + _per_token_group_quant_fp8[(triton.cdiv(M, groups_per_cta),)]( x, x_q, x_s, group_size, + M, N, eps, fp8_min=fp8_min, @@ -160,6 +167,7 @@ def lightllm_per_token_group_quant_fp8( NEED_MASK=BLOCK != group_size, USE_UE8M0_SCALE=use_ue8m0_scales, COPY_TOPK=copy_topk, + GROUPS_PER_CTA=groups_per_cta, num_warps=num_warps, num_stages=num_stages, ) diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 1217a2bc74..85257288f4 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -30,10 +30,10 @@ def quantize(self, weight: torch.Tensor, output: WeightPack): def finalize_moe_weight(self, moe_weight): from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - prepare_mega_moe_weights, + prepare_ep_moe_weights, ) - prepare_mega_moe_weights(moe_weight.w13, moe_weight.w2, self) + prepare_ep_moe_weights(moe_weight.w13, moe_weight.w2, self) def apply( self, diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 5b9e64ac58..1376e4ab6a 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -171,6 +171,8 @@ def __init__(self): self.ep_mega_moe_buffer = None self.ep_mega_moe_mma_type = None self.ep_mega_moe_quant_method = None + self.ep_triton_moe_buffer = None + self.ep_triton_moe_quant_method = None self.ep_num_sms = None def __len__(self): @@ -239,8 +241,8 @@ def new_deepep_group( ``expert_quant_method_names`` 是各 MoE 层最终绑定的 quant method 名称集合。 同一个模型可能逐层混用多种 expert quant method:满足约束的 SM100 FP4 和 - SM90 FP8 层走 Mega MoE,其他层走 DeepEP legacy 路径。这里只为实际存在的 - 执行路径分配 buffer,避免为未使用的路径长期占用显存。 + SM90 FP8 层走 Mega/Triton EP MoE,其他层走 DeepEP legacy 路径。这里只为实际 + 存在的执行路径分配 buffer,避免为未使用的路径长期占用显存。 """ args = get_env_start_args() enable_ep_moe = args.enable_ep_moe @@ -255,6 +257,8 @@ def new_deepep_group( self.ep_mega_moe_buffer = None self.ep_mega_moe_mma_type = None self.ep_mega_moe_quant_method = None + self.ep_triton_moe_buffer = None + self.ep_triton_moe_quant_method = None self.ep_num_sms = None return assert HAS_DEEPEP, "deep_ep is required for expert parallelism" @@ -289,6 +293,7 @@ def new_deepep_group( allow_multiple_reduction=True, ) self.ep_mega_moe_buffer = None + self.ep_triton_moe_buffer = None self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None @@ -297,42 +302,33 @@ def new_deepep_group( self.ep_mega_moe_mma_type = None self.ep_mega_moe_quant_method = None - if is_sm100_gpu() and FP4_MOE_QUANT_METHOD in expert_quant_method_names: + self.ep_triton_moe_quant_method = None + if args.ep_moe_backend == "triton" and FP8_MOE_QUANT_METHOD in expert_quant_method_names: + self.ep_triton_moe_quant_method = FP8_MOE_QUANT_METHOD + + if ( + self.ep_triton_moe_quant_method is None + and is_sm100_gpu() + and FP4_MOE_QUANT_METHOD in expert_quant_method_names + ): self.ep_mega_moe_mma_type = "fp8xfp4" self.ep_mega_moe_quant_method = FP4_MOE_QUANT_METHOD - elif ( + elif self.ep_triton_moe_quant_method is None and ( enable_env_vars("LIGHTLLM_ENABLE_SM90_FP8_MEGA_MOE") and is_sm90_gpu() and FP8_MOE_QUANT_METHOD in expert_quant_method_names and total_redundant_experts == 0 + and args.nnodes == 1 + and not args.enable_rl ): self.ep_mega_moe_mma_type = "fp8xfp8" self.ep_mega_moe_quant_method = FP8_MOE_QUANT_METHOD - if self.ep_mega_moe_mma_type == "fp8xfp8": - import deep_gemm - - fallback_reason = None - if not hasattr(deep_gemm, "fp8_fp8_mega_moe") or not hasattr( - getattr(deep_gemm, "_C", None), "fp8_fp8_mega_moe" - ): - fallback_reason = ( - "the loaded DeepGEMM Python package and extension do not both provide fp8_fp8_mega_moe " - f"({getattr(deep_gemm, '__file__', '')})" - ) - elif getattr(args, "nnodes", 1) != 1: - fallback_reason = "Mega MoE only supports a single-node expert-parallel group" - elif getattr(args, "enable_rl", False): - fallback_reason = "online expert-weight updates require the canonical non-interleaved layout" - elif not has_nvlink(): - fallback_reason = "NVLink is unavailable" - if fallback_reason is not None: - logger.warning("Disable SM90 FP8 Mega MoE and use legacy DeepEP because %s", fallback_reason) - self.ep_mega_moe_mma_type = None - self.ep_mega_moe_quant_method = None enable_mega_moe_buffer = self.ep_mega_moe_mma_type is not None - has_legacy_moe_layer = not enable_mega_moe_buffer or any( - method_name != self.ep_mega_moe_quant_method for method_name in expert_quant_method_names + enable_triton_ep_moe_buffer = self.ep_triton_moe_quant_method is not None + fused_moe_quant_method = self.ep_triton_moe_quant_method or self.ep_mega_moe_quant_method + has_legacy_moe_layer = fused_moe_quant_method is None or any( + method_name != fused_moe_quant_method for method_name in expert_quant_method_names ) enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" @@ -387,23 +383,43 @@ def new_deepep_group( ) theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_prefill_num_experts, num_experts_per_tok) - deepep_sms = 0 if self.ep_mega_moe_mma_type == "fp8xfp8" and not has_legacy_moe_layer else theoretical_sms + use_all_sms_for_fp8 = enable_triton_ep_moe_buffer or self.ep_mega_moe_mma_type == "fp8xfp8" + deepep_sms = 0 if use_all_sms_for_fp8 and not has_legacy_moe_layer else theoretical_sms low_latency_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_decode_num_experts, num_experts_per_tok) self._set_num_sms_for_deep_gemm(deepep_sms, low_latency_sms) + if enable_triton_ep_moe_buffer: + from lightllm.common.basemodel.triton_kernel.fused_moe.sm90_fp8_triton_ep_moe import ( + SM90FP8TritonEPMoEBuffer, + ) + + self.ep_triton_moe_buffer = SM90FP8TritonEPMoEBuffer( + deepep_group, + num_experts=self.ll_decode_num_experts, + num_max_tokens_per_rank=self.ll_num_tokens, + topk=num_experts_per_tok, + hidden_size=self.ll_hidden, + intermediate_size=moe_intermediate_size, + ) + if enable_mega_moe_buffer: if moe_intermediate_size is None: raise ValueError("Mega MoE requires moe_intermediate_size or intermediate_size in model config") import deep_gemm + num_max_tokens_per_rank = ( + self.ll_decode_num_tokens + if self.ep_mega_moe_mma_type == "fp8xfp8" and args.run_mode == "decode" + else self.ll_num_tokens + ) mega_buffer_kwargs = ( {"mma_type": self.ep_mega_moe_mma_type} if self.ep_mega_moe_mma_type == "fp8xfp8" else {} ) self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, self.ll_decode_num_experts, - self.ll_num_tokens, + num_max_tokens_per_rank, num_experts_per_tok, self.ll_hidden, moe_intermediate_size, @@ -411,12 +427,13 @@ def new_deepep_group( ) logger.info( "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " - "mega_moe=%s, mega_moe_mma_type=%s, ll_prefill_num_experts=%s, " + "mega_moe=%s, mega_moe_mma_type=%s, triton_ep_moe=%s, ll_prefill_num_experts=%s, " "ll_decode_num_experts=%s, expert_quant_method_names=%s", enable_low_latency_buffer, self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, self.ep_mega_moe_mma_type, + enable_triton_ep_moe_buffer, self.ll_prefill_num_experts, self.ll_decode_num_experts, sorted(expert_quant_method_names), diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index b2cbb737bf..5464c3adcb 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -7,9 +7,8 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.deepseek2.infer_struct import Deepseek2InferStateInfo -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - use_mega_moe, -) +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP from functools import partial from lightllm.models.llama.yarn_rotary_utils import get_deepseek_mscale from lightllm.utils.envs_utils import get_env_start_args @@ -300,7 +299,11 @@ def overlap_tpsp_token_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): + if ( + not self.is_moe + or use_mega_moe(layer_weight.experts.quant_method) + or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) + ): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -426,7 +429,11 @@ def overlap_tpsp_context_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): + if ( + not self.is_moe + or use_mega_moe(layer_weight.experts.quant_method) + or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) + ): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 6ac1e6af02..08a07b10c4 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -3,6 +3,7 @@ from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer @@ -150,7 +151,11 @@ def overlap_tpsp_context_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if not self.enable_ep_moe or use_mega_moe(experts.quant_method): + if ( + not self.enable_ep_moe + or use_mega_moe(experts.quant_method) + or isinstance(experts.fuse_moe_impl, FuseMoeTritonEP) + ): input_embdings = self.context_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.context_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 @@ -254,7 +259,11 @@ def overlap_tpsp_token_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if not self.enable_ep_moe or use_mega_moe(experts.quant_method): + if ( + not self.enable_ep_moe + or use_mega_moe(experts.quant_method) + or isinstance(experts.fuse_moe_impl, FuseMoeTritonEP) + ): input_embdings = self.token_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.token_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 diff --git a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py index aa45940440..d5347857b5 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -6,9 +6,8 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - use_mega_moe, -) +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.envs_utils import get_env_start_args @@ -138,7 +137,11 @@ def overlap_tpsp_token_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): + if ( + not self.is_moe + or use_mega_moe(layer_weight.experts.quant_method) + or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) + ): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -250,7 +253,11 @@ def overlap_tpsp_context_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if not self.is_moe or use_mega_moe(layer_weight.experts.quant_method): + if ( + not self.is_moe + or use_mega_moe(layer_weight.experts.quant_method) + or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) + ): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index f43d8fe8b7..ddda140f4a 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -749,6 +749,17 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: action="store_true", help="""Whether to enable ep moe for deepseekv3 model.""", ) + parser.add_argument( + "--ep_moe_backend", + type=str, + choices=["auto", "triton"], + default="auto", + help=( + "EP MoE execution backend. 'auto' keeps the existing backend selection; " + "'triton' selects the single-node SM90 FP8 Triton peer backend on Prefill nodes " + "for experts resolved by --expert_dtype fp8." + ), + ) parser.add_argument( "--disable_ep_balance_monitor", action="store_true", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 77c1e853ae..5741dbe9ab 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -27,7 +27,7 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.utils.device_utils import is_sm100_gpu +from lightllm.utils.device_utils import is_sm100_gpu, is_sm90_gpu from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup logger = init_logger(__name__) @@ -179,6 +179,13 @@ def _launch_subprocesses(args: StartArgs): if args.enable_dp_prefill_balance: assert args.enable_tpsp_mix_mode and args.dp > 1, "need set --enable_tpsp_mix_mode firstly and --dp > 1" + if args.ep_moe_backend == "triton": + assert args.enable_ep_moe, "--ep_moe_backend triton requires --enable_ep_moe" + assert args.run_mode == "prefill", "--ep_moe_backend triton only supports --run_mode prefill" + assert args.nnodes == 1, "--ep_moe_backend triton only supports a single node" + assert not args.enable_prefill_eplb, "--ep_moe_backend triton does not support --enable_prefill_eplb" + assert is_sm90_gpu(), "--ep_moe_backend triton only supports SM90 GPUs" + if args.enable_prefill_eplb: assert args.enable_ep_moe, "--enable_prefill_eplb requires --enable_ep_moe" assert not args.enable_prefill_cudagraph, "--enable_prefill_eplb does not support --enable_prefill_cudagraph" diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 3d036f55ad..a742e694b4 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -182,6 +182,7 @@ class StartArgs: default="gpu_counter", metadata={"choices": ["cpu_counter", "pin_mem_counter", "gpu_counter"]} ) enable_ep_moe: bool = field(default=False) + ep_moe_backend: str = field(default="auto", metadata={"choices": ["auto", "triton"]}) disable_ep_balance_monitor: bool = field(default=False) enable_prefill_eplb: bool = field(default=False) eplb_num_redundant_experts_per_rank: int = field(default=2) diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py index cf1074730a..639457fbf5 100644 --- a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py +++ b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py @@ -35,6 +35,7 @@ def should_enable_ep_balance_monitor(args) -> bool: args.enable_prefill_cudagraph or is_sm100_gpu() or getattr(dist_group_manager, "ep_mega_moe_buffer", None) is not None + or getattr(dist_group_manager, "ep_triton_moe_buffer", None) is not None ): return False return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" From 9bb68ebabfb26a0579b8764155bcb5ddcd26380b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 16 Sep 2026 09:13:35 +0000 Subject: [PATCH 184/214] fix abort --- .../httpserver_for_pd_master/manager.py | 28 +++++++++++++------ 1 file changed, 19 insertions(+), 9 deletions(-) diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 582caa138a..19f91d1f77 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -58,6 +58,7 @@ def __init__( self.latest_success_infer_time = time.time() self.running_request_count = 0 self.next_request_queue_metric_time = 0.0 + self._abort_notify_tasks = set() # 高优先级请求仍可比普通请求等待更久,但通过请求参数向开启本地限流的 # P/D 节点传递有限的等待时间,避免资源异常时永久占用请求链路。 self.pd_high_priority_request_time_out_seconds = get_pd_high_priority_request_timeout_seconds() @@ -678,15 +679,24 @@ async def abort( except: pass - try: - await p_node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) - except: - pass - - try: - await d_node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) - except: - pass + async def notify_node(node: Optional[PD_Client_Obj]): + if node is None: + return + try: + await node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) + except BaseException: + pass + + # HTTP request cancellation must not cancel an ABORT while it is waiting for + # the node's websocket send lock. Keep independent tasks alive until both + # nodes have received the cleanup message or their connections fail. + notify_tasks = [asyncio.create_task(notify_node(node)) for node in (p_node, d_node) if node is not None] + self._abort_notify_tasks.update(notify_tasks) + for task in notify_tasks: + task.add_done_callback(self._abort_notify_tasks.discard) + + if notify_tasks: + await asyncio.gather(*(asyncio.shield(task) for task in notify_tasks)) return From 8ee409a5610d5a4a1023942a07ead1f020d269ec Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 17 Sep 2026 03:51:06 +0000 Subject: [PATCH 185/214] fix(api): ignore structured output constraints when xgrammar is disabled --- lightllm/server/core/objs/sampling_params.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index b874631a91..9fd03fa1b8 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -382,7 +382,8 @@ def init(self, tokenizer, **kwargs): guided_grammar = kwargs.get("guided_grammar", "") guided_json = kwargs.get("guided_json", "") if (guided_grammar or guided_json) and get_env_start_args().output_constraint_mode != "xgrammar": - raise ValueError("Structured output is not supported") + guided_grammar = "" + guided_json = "" self.guided_grammar = GuidedGrammar() self.guided_grammar.initialize(guided_grammar, tokenizer) From 2bab226a39a8de116c631b5fa664c661cc99053d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 18 Sep 2026 07:26:36 +0000 Subject: [PATCH 186/214] fix(stream): emit keepalive deltas for long DSML tool arguments --- lightllm/server/function_call_parser.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index a4bba104af..eb217d6b91 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1686,6 +1686,12 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami continue if self.current_tool_name_sent and not has_new_param_end: + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + parameters="", + ) + ) return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) partial_match = self.partial_invoke_regex.match(current_text) From 2483650c972013f6044b1e4d5d78fbe9f3cfdcaf Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Fri, 18 Sep 2026 09:26:08 +0000 Subject: [PATCH 187/214] add LIGHTLLM_DSV4_THINKING_EFFORT --- lightllm/models/deepseek_v4/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 09ed7a60ef..57ca1215de 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -590,7 +590,7 @@ def apply_chat_template( thinking_mode = "thinking" if thinking else "chat" effort = kwargs.get("reasoning_effort") if thinking and effort is None: - effort = "high" + effort = os.getenv("LIGHTLLM_DSV4_THINKING_EFFORT", "high") if effort not in ("max", "high", None): effort = None encoding = self._get_encoding_module() From 21763ad6d084bc28f39d2984d410fbb31bfecc2f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 06:00:43 +0000 Subject: [PATCH 188/214] support tp and tppd --- .../deepseek4_mem_manager.py | 25 +++++++++++-------- lightllm/server/api_start.py | 4 --- 2 files changed, 15 insertions(+), 14 deletions(-) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 78d92fe251..4cdc50e3fb 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1287,8 +1287,10 @@ def write_mem_to_page_kv_move_buffer( mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) from lightllm.models.deepseek_v4.triton_kernel.pd_cache_io import pack_pd_cache_page + dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] + assert len(dp_mems) == dp_world_size return pack_pd_cache_page( - mem_managers[dp_index], + dp_mems[0], self.pd_cache_layout, mem_indexes_gpu, self.kv_move_buffer[page_index], @@ -1315,13 +1317,16 @@ def read_page_kv_move_buffer_to_mem( mem_indexes_gpu = pin_mem_indexes.cuda(non_blocking=True) from lightllm.models.deepseek_v4.triton_kernel.pd_cache_io import unpack_pd_cache_page - unpack_pd_cache_page( - mem_managers[dp_index], - self.pd_cache_layout, - mem_indexes_gpu, - self.kv_move_buffer[page_index], - start_kv_index, - request_kv_len, - req_idx, - ) + dp_mems = mem_managers[(dp_index * dp_world_size) : ((dp_index + 1) * dp_world_size)] + assert len(dp_mems) == dp_world_size + for mem_manager in dp_mems: + unpack_pd_cache_page( + mem_manager, + self.pd_cache_layout, + mem_indexes_gpu, + self.kv_move_buffer[page_index], + start_kv_index, + request_kv_len, + req_idx, + ) return diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 5741dbe9ab..3655ee1c12 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -67,10 +67,6 @@ def _launch_subprocesses(args: StartArgs): if args.run_mode not in ["normal", "prefill", "decode", "visual_only"]: return - if args.run_mode in ("prefill", "decode") and model_type == "deepseek_v4": - if args.tp != args.dp: - raise ValueError("DeepSeek-V4 PD requires one TP rank per DP replica (--tp must equal --dp)") - # 通过模型的参数判断是否是多模态模型,包含哪几种模态, 并设置是否启动相应得模块 if args.disable_vision is None: if has_vision_module(args.model_dir): From 8ce69ccb0e60b2f0f2e9366d7848a58ac71f2f8f Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 06:01:43 +0000 Subject: [PATCH 189/214] perf(dspark): merge context WKV projections across draft stages --- .../models/deepseek_v4_dspark/infer_struct.py | 3 +++ .../layer_infer/transformer_layer_infer.py | 12 +++++---- lightllm/models/deepseek_v4_dspark/model.py | 26 ++++++++++++++++++- 3 files changed, 35 insertions(+), 6 deletions(-) diff --git a/lightllm/models/deepseek_v4_dspark/infer_struct.py b/lightllm/models/deepseek_v4_dspark/infer_struct.py index 14837466f8..175fc46163 100644 --- a/lightllm/models/deepseek_v4_dspark/infer_struct.py +++ b/lightllm/models/deepseek_v4_dspark/infer_struct.py @@ -10,6 +10,9 @@ class DeepseekV4DSparkInferStateInfo(DeepseekV4InferStateInfo): """DeepSeek-V4 metadata with non-causal visibility across one DSpark block.""" + # Produced by the first context layer after DP rebalance, consumed per stage. + context_kv: tuple + def init_some_extra_state(self, model): super().init_some_extra_state(model) if self.is_prefill or self.mtp_draft_swa_pages is None: diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py index 72d08bb720..d8d5a30ffe 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py @@ -1,6 +1,6 @@ import torch -from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo +from lightllm.models.deepseek_v4_dspark.infer_struct import DeepseekV4DSparkInferStateInfo from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import ( DeepseekV4TransformerLayerInfer, ) @@ -15,22 +15,24 @@ class DeepseekV4DSparkTransformerLayerInfer(DeepseekV4TransformerLayerInfer): def __init__(self, layer_num, network_config): super().__init__(layer_num, network_config) final_layer = network_config["n_layer"] + network_config["dspark_layer_num"] - 1 + self.stage_id = layer_num - network_config["n_layer"] self.is_last_layer = layer_num == final_layer assert self.compress_ratio == 0, "DeepSeek-V4 DSpark draft layers must be SWA-only" def context_forward( self, input_embdings: torch.Tensor, - infer_state: DeepseekV4InferStateInfo, + infer_state: DeepseekV4DSparkInferStateInfo, layer_weight: DeepseekV4DSparkTransformerLayerWeight, ) -> torch.Tensor: """Write target hidden rows into this stage without running draft attention/FFN.""" - full_input = self._tpsp_allgather(input=input_embdings, infer_state=infer_state) - qkv = layer_weight.wq_a_wkv_.mm(full_input, use_custom_tensor_mananger=False) + if self.stage_id == 0: + all_kv = self.context_wkv_weight.mm(input_embdings, use_custom_tensor_mananger=False) + infer_state.context_kv = all_kv.split(self.head_dim_, dim=-1) infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, mem_index=infer_state.mem_index, - kv=qkv[:, -self.head_dim_ :], + kv=infer_state.context_kv[self.stage_id], kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, freqs_cis=self.freqs_cis, diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 458f162b88..1138388a0a 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -10,6 +10,7 @@ import lightllm.utils.petrel_helper as utils from lightllm.common.basemodel import TpPartBaseModel from lightllm.common.basemodel.batch_objs import ModelInput, ModelOutput +from lightllm.common.basemodel.layer_weights.meta_weights import ROWMMWeight from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel from lightllm.models.deepseek_v4_dspark.infer_struct import ( DeepseekV4DSparkInferStateInfo, @@ -110,6 +111,15 @@ def _init_weights(self, start_layer_index=None): ) for stage_id in range(self.config["dspark_layer_num"]) ] + self.context_wkv_weight = ROWMMWeight( + in_dim=self.config["hidden_size"], + out_dims=[layer.head_dim for layer in self.trans_layers_weight], + weight_names=[layer.wq_a_wkv_.weight_names[1] for layer in self.trans_layers_weight], + data_type=self.data_type, + quant_method=self.trans_layers_weight[0].get_quant_method("wq_a"), + tp_rank=0, + tp_world_size=1, + ) def _init_req_manager(self): self.req_manager = self.main_model.req_manager @@ -173,6 +183,7 @@ def _init_custom(self): layer.freqs_cis = self._freqs_cis_sliding layer.cos_compress_table = self._cos_cached_compress layer.sin_compress_table = self._sin_cached_compress + self.layers_infer[0].context_wkv_weight = self.context_wkv_weight def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: if not model_input.is_prefill: @@ -228,7 +239,9 @@ def _init_prefill_cuda_graph(self): def load_weights(self, weight_dict: dict): if weight_dict: - return super().load_weights(weight_dict) + super().load_weights(weight_dict) + self._load_context_wkv_weights() + return index_file = os.path.join(self.weight_dir_, "model.safetensors.index.json") assert utils.PetrelHelper.exists(index_file), "DeepSeek-V4 DSpark requires model.safetensors.index.json" @@ -254,4 +267,15 @@ def load_weights(self, weight_dict: dict): self.pre_post_weight.verify_load() [weight.verify_load() for weight in self.trans_layers_weight] + self._load_context_wkv_weights() logger.info("loaded DeepSeek-V4 DSpark weights: %d tensors", loaded_key_count) + + def _load_context_wkv_weights(self): + for layer, dest in zip(self.trans_layers_weight, self.context_wkv_weight.mm_param_list): + source = layer.wq_a_wkv_.mm_param_list[1] + for name in ("weight", "weight_scale", "weight_zero_point"): + dest_tensor = getattr(dest, name) + if dest_tensor is not None: + dest_tensor.copy_(getattr(source, name)) + dest.load_ok[:] = source.load_ok + assert self.context_wkv_weight.verify_load(), "Loading DSpark context WKV projection failed" From 8415357e6b6fd5048a598970238c014e2508947c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 07:04:09 +0000 Subject: [PATCH 190/214] fix(dsv4): fall back to FlashMLA when workspace op is unavailable --- .../attention/nsa/dsv4_fp8_flashmla_sparse.py | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py index 2205439b55..45d2760406 100644 --- a/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py +++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py @@ -4,11 +4,17 @@ import torch from vllm.v1.attention.ops import flashmla +from lightllm.utils.log_utils import init_logger + from ..base_att import AttControl, BaseAttBackend, BaseDecodeAttState, BasePrefillAttState if TYPE_CHECKING: from lightllm.common.basemodel.infer_struct import InferStateInfo +logger = init_logger(__name__) +if not hasattr(torch.ops._flashmla_C, "sparse_decode_fwd_with_workspace"): + logger.warning("FlashMLA sparse_decode_fwd_with_workspace is unavailable; prefill uses flash_mla_with_kvcache") + # The current FlashMLA MODEL1 binary only instantiates these Q-head counts. _SUPPORTED_Q_HEADS = (64, 128) @@ -28,12 +34,9 @@ def _view_cache(buffer: torch.Tensor, page_size: int) -> torch.Tensor: return buffer[:, :byte_num].view(buffer.shape[0], page_size, 1, DSV4_MLA_BYTES_PER_TOKEN) -def _flashmla_sparse_decode_with_workspace(kwargs: dict, sched_meta, o_accum: torch.Tensor, lse_accum: torch.Tensor): - try: - op = torch.ops._flashmla_C.sparse_decode_fwd_with_workspace - except AttributeError as exc: - raise RuntimeError("DeepSeek-V4 prefill requires the workspace-enabled vllm._flashmla_C extension") from exc - +def _flashmla_sparse_decode_with_workspace( + op, kwargs: dict, sched_meta, o_accum: torch.Tensor, lse_accum: torch.Tensor +): out, lse, new_tile_scheduler_metadata, new_num_splits = op( kwargs["q"], kwargs["k_cache"], @@ -62,6 +65,7 @@ def __init__(self, model): self.real_q_head_num = model.config["num_attention_heads"] // model.tp_world_size_ self.padded_q_head_num = get_dsv4_flashmla_padded_q_heads(self.real_q_head_num) self.compress_ratios = tuple(dict.fromkeys(model.config["compress_ratios"])) + self._flashmla_workspace_op = getattr(torch.ops._flashmla_C, "sparse_decode_fwd_with_workspace", None) def _flashmla_att( self, @@ -112,10 +116,11 @@ def _flashmla_att( ) if flashmla_out is not None: kwargs["out"] = flashmla_out - if flashmla_o_accum is None: + if flashmla_o_accum is None or self._flashmla_workspace_op is None: full_out, _ = flashmla.flash_mla_with_kvcache(**kwargs) else: full_out, _ = _flashmla_sparse_decode_with_workspace( + self._flashmla_workspace_op, kwargs, sched_meta, flashmla_o_accum, From 16176f89e7ebcf9e08cf3f15d3b0a567d08f5105 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 18:00:48 +0800 Subject: [PATCH 191/214] Unify DeepSeek-V4 compressed history on 256-token pages --- lightllm/common/basemodel/batch_objs.py | 39 +-- .../deepseek4_mem_manager.py | 185 ++--------- .../kv_cache_mem_manager/operator/deepseek.py | 12 + lightllm/common/req_manager/deepseek4.py | 291 ++---------------- lightllm/models/deepseek_v4/infer_struct.py | 1 - .../deepseek_v4/layer_infer/compressor.py | 11 +- .../layer_infer/transformer_layer_infer.py | 1 - lightllm/models/deepseek_v4/model.py | 76 ++--- .../build_compress_index_dsv4.py | 16 +- .../triton_kernel/cache_staging_io.py | 10 +- .../deepseek_v4/triton_kernel/cpu_cache_io.py | 8 +- .../destindex_copy_indexer_k_dsv4.py | 10 +- .../deepseek_v4/triton_kernel/dp_cache_io.py | 32 +- .../triton_kernel/gather_c4_indexer_k_dsv4.py | 4 +- .../deepseek_v4/triton_kernel/pd_cache_io.py | 2 - .../layer_infer/post_layer_infer.py | 6 +- lightllm/models/deepseek_v4_dspark/model.py | 34 +- lightllm/server/api_start.py | 2 + .../router/dynamic_prompt/radix_cache.py | 38 --- .../server/router/model_infer/infer_batch.py | 94 ++---- .../model_infer/mode_backend/base_backend.py | 30 +- .../dp_backend/dp_shared_kv_trans.py | 18 +- .../mode_backend/dsv4_multi_level_kv_cache.py | 11 +- .../pd/decode_node_impl/decode_impl.py | 14 +- .../dp_overlap_proposers/eagle_with_att.py | 3 - .../proposers/eagle_with_att.py | 1 - .../static_inference/static_benchmark.py | 11 +- test/unit/test_deepseek_v4_dspark.py | 75 ++--- .../common/test_deepseek4_paged_cache.py | 192 ++++++++++++ 29 files changed, 393 insertions(+), 834 deletions(-) create mode 100644 unit_tests/common/test_deepseek4_paged_cache.py diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 05f3af23a7..5f2f3a587c 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -41,10 +41,10 @@ class ModelInput: b_prefill_start_loc: torch.Tensor = None multimodal_params: list = None # cpu 变量 - mem_indexes_cpu: torch.Tensor = None b_req_idx_cpu: torch.Tensor = None b_mtp_index_cpu: torch.Tensor = None b_seq_len_cpu: torch.Tensor = None + b_ready_cache_len_cpu: torch.Tensor = None # prefill 阶段使用的参数,但是不是推理过程使用的参数,是推理外部进行资源管理 # 的一些变量 # 标记 prefill 请求是否会在本轮产生输出。Prefill 必填(空 batch 使用空 list),decode 不使用。 @@ -60,8 +60,6 @@ class ModelInput: # 回收;GPU tensor 供 attention 直接计算 block 的物理 SWA 槽。 mtp_draft_swa_pages_cpu: Optional[torch.Tensor] = None mtp_draft_swa_pages: Optional[torch.Tensor] = None - # 主模型为 None: 准备所有 MTP 列;draft 首轮为 (): 无新槽;draft 追加后为 (k,): 只准备新槽。 - mtp_decode_slot_prepare_indices: Optional[tuple] = None def _capture_cpu_mirror(self, tensor_name: str, mirror_name: str): tensor = getattr(self, tensor_name) @@ -73,40 +71,7 @@ def capture_cpu_mirrors(self): self._capture_cpu_mirror("b_req_idx", "b_req_idx_cpu") self._capture_cpu_mirror("b_mtp_index", "b_mtp_index_cpu") self._capture_cpu_mirror("b_seq_len", "b_seq_len_cpu") - return - - def make_mtp_draft_input(self): - model_input = copy.copy(self) - model_input.b_seq_len = self.b_seq_len.clone() - model_input.b_seq_len_cpu = self.b_seq_len_cpu.clone() - model_input.mtp_decode_slot_prepare_indices = () - return model_input - - def advance_mtp_decode_step( - self, - new_mem_indexes_cpu: torch.Tensor, - new_mem_indexes: torch.Tensor, - max_mtp_index: int, - ): - self.b_seq_len += 1 - self.b_seq_len_cpu += 1 - self.max_kv_seq_len += 1 - self.mtp_decode_slot_prepare_indices = (max_mtp_index,) - slots_per_req = max_mtp_index + 1 - self.mem_indexes_cpu = torch.cat( - [ - self.mem_indexes_cpu.view(-1, slots_per_req)[:, 1:], - new_mem_indexes_cpu.view(-1, 1), - ], - dim=1, - ).view(-1) - self.mem_indexes = torch.cat( - [ - self.mem_indexes.view(-1, slots_per_req)[:, 1:], - new_mem_indexes.view(-1, 1), - ], - dim=1, - ).view(-1) + self._capture_cpu_mirror("b_ready_cache_len", "b_ready_cache_len_cpu") return def to_cuda(self): diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 4cdc50e3fb..5d38884ea9 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -302,10 +302,6 @@ class DeepseekV4CpuCacheLoadPlan: resume_full_slots_long: torch.Tensor resume_swa_pages: torch.Tensor resume_swa_page_deltas: torch.Tensor - history_c4_full_slots_long: Optional[torch.Tensor] - history_c4_pages: Optional[torch.Tensor] - history_c4_page_deltas: Optional[torch.Tensor] - history_c128_full_slots_long: Optional[torch.Tensor] class PackedPagePool: @@ -389,8 +385,8 @@ class DeepseekV4MemoryManager(MemoryManager): 页 allocator 触底时先走 swa free hook(radix 对 ref==0 节点 free)再 assert。 没有 ring buffer,prefill chunk 大小不受 sliding_window 限制。 - ``c4_pool``/``c128_pool``: 压缩 latent,按 qwen3next 的层号压实手法只为压缩层建层; - c4 另带 packed indexer-K 池。槽位映射(``full_to_c4/c128_indexs``)以组末 token 的 full - 槽位为键(prep 阶段分配/scatter),``free`` 级联回收,与 swa 完全同构。 + c4 另带 packed indexer-K 池。统一 256-token 页拥有三个 packed 历史页; + 压缩槽位直接由组末 full slot 除以压缩率得到,随 full 页分配和释放。 - 写入走标准 operator 路径(``pack_mla_kv_to_cache``),内部为 triton packed writer; torch codecs 保留为 ABI 的可执行规格(单测 oracle)。 """ @@ -423,6 +419,8 @@ def __init__( mem_fraction=0.9, memory_reservations=None, ): + if get_env_start_args().page_size != DSV4_PROMPT_CACHE_PAGE_SIZE: + raise ValueError("DeepSeek-V4 requires --page_size 256") assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" assert head_dim == self.mla_head_dim, f"DeepSeek-V4 packed KV 期望 head_dim={self.mla_head_dim}" assert ( @@ -543,24 +541,20 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): # swa free hook(可选): 页 allocator 触底时回调(radix 对 ref==0 节点 free swa 页), # 由 backend 在 radix cache 创建后 register;assert 仍是最后防线。 self._free_radix_unreferenced_swa_fn = None - self.full_to_swa_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") - self.full_to_swa_indexs[size] = self.swa_pool.HOLD_TOKEN_MEMINDEX + self.HOLD_TOKEN_MEMINDEX = size + self.full_to_swa_indexs = torch.full((size + self.page_size,), -1, dtype=torch.int32, device="cuda") + self.full_to_swa_indexs[size:] = ( + self.swa_size + torch.arange(self.page_size, device="cuda") % DSV4_SWA_PAGE_SIZE + ) self.c4_size = _ceil_div(size, 4) self.c128_size = _ceil_div(size, 128) self.c4_pool: Optional[PackedPagePool] = None self.c4_indexer_pool: Optional[PackedPagePool] = None - self.c4_page_allocator: Optional[KvCacheAllocator] = None - self.c4_page_live_count: Optional[torch.Tensor] = None self.c128_pool: Optional[PackedPagePool] = None - self.c128_allocator: Optional[KvCacheAllocator] = None self.c4_state_buffer: Optional[torch.Tensor] = None self.c4_indexer_state_buffer: Optional[torch.Tensor] = None self.c128_state_buffer: Optional[torch.Tensor] = None - # 压缩槽映射: 键 = 组末 token(位置 (g+1)%ratio==0)的 full 槽位,值 = 压缩池槽位。 - # 与 full_to_swa_indexs 同构: radix 持有 full 槽 => 映射行存活,free 级联回收。 - self.full_to_c4_indexs: Optional[torch.Tensor] = None - self.full_to_c128_indexs: Optional[torch.Tensor] = None if self.n_c4 > 0: self.c4_pool = PackedPagePool( size=self.c4_size, @@ -577,14 +571,6 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): data_bytes=self.indexer_head_dim, scale_bytes=DSV4_INDEXER_SCALE_BYTES, ) - self.c4_num_pages = self.c4_size // DSV4_C4_PAGE_SIZE - assert self.c4_num_pages > 0, "DeepSeek-V4 c4 pool must have at least one usable full page" - self.c4_page_allocator = KvCacheAllocator( - self.c4_num_pages, shared_name=f"{server}_dsv4_c4_can_use_page_num_{rank_in_node}" - ) - self.c4_page_live_count = torch.zeros((self.c4_pool.num_pages,), dtype=torch.int32, device="cuda") - self.full_to_c4_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") - self.full_to_c4_indexs[size] = self.c4_pool.HOLD_TOKEN_MEMINDEX # c4 compressor 在途状态(attention + indexer): swa 页派生寻址(翻译③),随 swa 页 # 生灭 -> radix 命中零拷贝续算。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), # 取整到 ratio;末行哨兵 kv=0/score=-inf(KVAndScore.clear 语义),其余行由内核在 @@ -607,11 +593,6 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): scale_bytes=self.mla_scale_bytes, align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES, ) - self.c128_allocator = KvCacheAllocator( - self.c128_size, shared_name=f"{server}_dsv4_c128_can_use_token_num_{rank_in_node}" - ) - self.full_to_c128_indexs = torch.full((size + 1,), -1, dtype=torch.int32, device="cuda") - self.full_to_c128_indexs[size] = self.c128_pool.HOLD_TOKEN_MEMINDEX # c128 compressor 在途状态按 request 寻址。每个 request 保留完整的 128-token # 聚合窗口以及 MTP 候选余量;最后一行供无效请求/位置读取哨兵。 state_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 @@ -674,8 +655,6 @@ def get_loadable_cpu_cache_end( requested_end: int, full_token_capacity: int, swa_page_capacity: int, - c4_page_capacity: int, - c128_slot_capacity: int, ) -> int: """Return the farthest loadable checkpoint boundary, or zero. @@ -697,10 +676,6 @@ def get_loadable_cpu_cache_end( return 0 token_capacity = int(full_token_capacity) - if self.n_c4: - token_capacity = min(token_capacity, int(c4_page_capacity) * DSV4_PROMPT_CACHE_PAGE_SIZE) - if self.n_c128: - token_capacity = min(token_capacity, int(c128_slot_capacity) * 128) token_capacity = token_capacity // DSV4_PROMPT_CACHE_PAGE_SIZE * DSV4_PROMPT_CACHE_PAGE_SIZE loadable_end = min(requested_end, (loaded_start + token_capacity) // page * page) return loadable_end if loadable_end > loaded_start else 0 @@ -730,8 +705,6 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C device = self.full_to_swa_indexs.device full_indexes_cpu = self.alloc(token_num) swa_pages_cpu = self._alloc_swa_pages(2) - c4_pages_cpu = self.alloc_c4_pages(block_num) if self.n_c4 else None - c128_slots_cpu = self.alloc_c128(block_num * 2) if self.n_c128 else None mem_indexes = full_indexes_cpu.to(device, non_blocking=True) history_full_slots = mem_indexes.view(block_num, DSV4_PROMPT_CACHE_PAGE_SIZE) @@ -743,17 +716,8 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C + torch.arange(DSV4_SWA_PAGE_SIZE, dtype=torch.int32, device=device)[None, :] ).reshape(-1) - history_c4_slots = None - if c4_pages_cpu is not None: - c4_pages = c4_pages_cpu.to(device, non_blocking=True) - history_c4_slots = ( - c4_pages[:, None] * DSV4_C4_PAGE_SIZE - + torch.arange(DSV4_C4_PAGE_SIZE, dtype=torch.int32, device=device)[None, :] - ) - - history_c128_slots = None - if c128_slots_cpu is not None: - history_c128_slots = c128_slots_cpu.to(device, non_blocking=True).view(block_num, 2) + history_c4_slots = history_full_slots[:, 3::4] // 4 if self.n_c4 else None + history_c128_slots = history_full_slots[:, 127::128] // 128 if self.n_c128 else None resume_full_slots_long = resume_full_slots.long() resume_swa_pages = swa_pages @@ -764,21 +728,6 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C device=device, ) - history_c4_full_slots_long = history_c4_pages = history_c4_page_deltas = None - if history_c4_slots is not None: - history_c4_full_slots_long = history_full_slots[:, 3::4].long() - history_c4_pages = c4_pages - history_c4_page_deltas = torch.full( - history_c4_pages.shape, - DSV4_C4_PAGE_SIZE, - dtype=torch.int32, - device=device, - ) - - history_c128_full_slots_long = None - if history_c128_slots is not None: - history_c128_full_slots_long = history_full_slots[:, 127::128].long() - return DeepseekV4CpuCacheLoadPlan( loaded_start=loaded_end - token_num, loaded_end=loaded_end, @@ -791,21 +740,12 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C resume_full_slots_long=resume_full_slots_long, resume_swa_pages=resume_swa_pages, resume_swa_page_deltas=resume_swa_page_deltas, - history_c4_full_slots_long=history_c4_full_slots_long, - history_c4_pages=history_c4_pages, - history_c4_page_deltas=history_c4_page_deltas, - history_c128_full_slots_long=history_c128_full_slots_long, ) def commit_cpu_cache_load_plan(self, plan: DeepseekV4CpuCacheLoadPlan) -> None: """Publish an unpacked plan without allocating at the transaction boundary.""" self.full_to_swa_indexs[plan.resume_full_slots_long] = plan.resume_swa_slots self.swa_page_live_count.index_add_(0, plan.resume_swa_pages, plan.resume_swa_page_deltas) - if plan.history_c4_slots is not None: - self.full_to_c4_indexs[plan.history_c4_full_slots_long] = plan.history_c4_slots - self.c4_page_live_count.index_add_(0, plan.history_c4_pages, plan.history_c4_page_deltas) - if plan.history_c128_slots is not None: - self.full_to_c128_indexs[plan.history_c128_full_slots_long] = plan.history_c128_slots return # ------------------------------------------------------------------ swa slot lifecycle @@ -929,23 +869,22 @@ def alloc_swa_decode( self._update_swa_page_counts(slots, 1) return - def alloc_dspark_swa_block(self, mem_indexes: torch.Tensor, block_size: int): + def alloc_dspark_swa_block(self, token_num: int, block_size: int): """Assign one temporary SWA scratch page to each DSpark proposal block. - Proposal KV is consumed by the draft model immediately and every full - slot is released after the following target verify. Keeping the whole + Proposal KV is consumed within the draft forward, then its scratch + pages are released. Request token pages stay reserved. Keeping the whole block in a private page avoids the host-side sequence-length decision required by the position-aligned target cache. Attention still uses the absolute positions stored in the request table; physical SWA slots only identify the packed KV rows. """ - mem_indexes = mem_indexes.reshape(-1) block_size = int(block_size) assert block_size > 0 assert block_size <= DSV4_SWA_PAGE_SIZE - assert mem_indexes.numel() % block_size == 0 + assert token_num % block_size == 0 - req_num = mem_indexes.numel() // block_size + req_num = token_num // block_size if req_num == 0: return ( torch.empty((0,), dtype=torch.int32, device="cpu"), @@ -959,17 +898,9 @@ def alloc_dspark_swa_block(self, mem_indexes: torch.Tensor, block_size: int): # 计算,不发布到 target 的全局 full->SWA 映射,也不参与 live count。 return pages_cpu, pages - def free_dspark_swa_block( - self, - mem_indexes_cpu: torch.Tensor, - pages_cpu: torch.Tensor, - ) -> None: - """Return DSpark scratch resources using their original CPU handles.""" - assert not mem_indexes_cpu.is_cuda - assert not pages_cpu.is_cuda + def free_dspark_swa_block(self, pages_cpu: torch.Tensor) -> None: + """Release only proposal scratch; token pages remain owned by the request.""" self.swa_page_allocator.free(pages_cpu) - MemoryManager.free(self, mem_indexes_cpu) - return def evict_swa(self, full_slots: torch.Tensor) -> None: """回收 full 槽位对应的 swa 槽(出窗惰性回收 / free 级联 / 压力阀共用)。 @@ -977,7 +908,7 @@ def evict_swa(self, full_slots: torch.Tensor) -> None: if full_slots.numel() == 0: return full_slots = full_slots.to(self.full_to_swa_indexs.device, non_blocking=True).reshape(-1) - full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) + full_slots = torch.unique(full_slots[full_slots < self.size]) if full_slots.numel() == 0: return swa_slots = self.full_to_swa_indexs[full_slots] @@ -992,62 +923,6 @@ def evict_swa(self, full_slots: torch.Tensor) -> None: self.swa_page_allocator.free(empty.to(torch.int32)) return - def _evict_compress(self, full_slots: torch.Tensor, mapping: torch.Tensor, allocator: KvCacheAllocator) -> None: - full_slots = full_slots.to(mapping.device, non_blocking=True).reshape(-1) - # 去重: 同批重复槽会 gather 出重复的压缩槽 -> allocator 双重释放(free 已去重,直呼叫方防御)。 - full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) - if full_slots.numel() == 0: - return - slots = mapping[full_slots] - valid = slots >= 0 - valid_slots = slots[valid] - if valid_slots.numel() == 0: - return - mapping[full_slots[valid]] = -1 - # The allocator's blocking D2H copy must also finish invalidating the - # old mapping before another stream can reuse the returned slots. - allocator.free(valid_slots) - return - - def alloc_c4_pages(self, need_pages: int) -> torch.Tensor: - assert self.c4_page_allocator is not None, "DeepSeek-V4 c4 page allocator is not initialized" - return self.c4_page_allocator.alloc(need_pages) - - def count_c4_slots(self, c4_slots: torch.Tensor, delta: int) -> torch.Tensor: - """按 c4 slot 所在页更新存活计数,返回逐 slot 的页号。""" - assert self.c4_page_live_count is not None, "DeepSeek-V4 c4 page live count is not initialized" - pages = torch.div(c4_slots, DSV4_C4_PAGE_SIZE, rounding_mode="floor") - ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) - self.c4_page_live_count.index_add_(0, pages, ones) - return pages - - def evict_c4(self, full_slots: torch.Tensor) -> None: - """回收 full 槽位(组末 token)映射的 c4 槽。非组末/未映射(-1)的槽位跳过。""" - if self.c4_page_allocator is None or full_slots.numel() == 0: - return - full_slots = full_slots.to(self.full_to_c4_indexs.device, non_blocking=True).reshape(-1) - full_slots = torch.unique(full_slots[full_slots != self.HOLD_TOKEN_MEMINDEX]) - if full_slots.numel() == 0: - return - slots = self.full_to_c4_indexs[full_slots] - valid = slots >= 0 - valid_slots = slots[valid] - if valid_slots.numel() == 0: - return - self.full_to_c4_indexs[full_slots[valid]] = -1 - touched = torch.unique(self.count_c4_slots(valid_slots, -1)) - empty = touched[self.c4_page_live_count[touched] == 0] - if empty.numel() > 0: - self.c4_page_allocator.free(empty.to(torch.int32)) - return - - def evict_c128(self, full_slots: torch.Tensor) -> None: - """回收 full 槽位(组末 token)映射的 c128 槽。非组末/未映射(-1)的槽位跳过。""" - if self.c128_allocator is None or full_slots.numel() == 0: - return - self._evict_compress(full_slots, self.full_to_c128_indexs, self.c128_allocator) - return - # ------------------------------------------------------------------ alloc/free (cascade) def free(self, free_index: Union[torch.Tensor, List[int]]) -> None: """释放 full token 槽位,级联回收其 swa 槽与 c4/c128 压缩槽。radix 驱逐、请求释放/暂停都走这里。 @@ -1058,8 +933,6 @@ def free(self, free_index: Union[torch.Tensor, List[int]]) -> None: if free_index.numel() > 0: free_index = torch.unique(free_index) self.evict_swa(free_index) - self.evict_c4(free_index) - self.evict_c128(free_index) super().free(free_index) return @@ -1068,24 +941,11 @@ def free_all(self): self.swa_page_allocator.free_all() self.swa_page_live_count.zero_() self.full_to_swa_indexs.fill_(-1) - self.full_to_swa_indexs[self.HOLD_TOKEN_MEMINDEX] = self.swa_pool.HOLD_TOKEN_MEMINDEX - if self.c4_page_allocator is not None: - self.c4_page_allocator.free_all() - self.c4_page_live_count.zero_() - self.full_to_c4_indexs.fill_(-1) - self.full_to_c4_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c4_pool.HOLD_TOKEN_MEMINDEX - if self.c128_allocator is not None: - self.c128_allocator.free_all() - self.full_to_c128_indexs.fill_(-1) - self.full_to_c128_indexs[self.HOLD_TOKEN_MEMINDEX] = self.c128_pool.HOLD_TOKEN_MEMINDEX + self.full_to_swa_indexs[self.size :] = ( + self.swa_size + torch.arange(self.page_size, device=self.full_to_swa_indexs.device) % DSV4_SWA_PAGE_SIZE + ) return - def alloc_c128(self, need_size) -> torch.Tensor: - return self.c128_allocator.alloc(need_size) - - def free_c128(self, free_index) -> None: - self.c128_allocator.free(free_index) - # ------------------------------------------------------------------ packed codecs (torch reference) # 与 sglang/vllm 的 fp8_ds_mla 字节布局逐位对齐(ue8m0 幂次 scale)。这些 torch 实现是该 ABI 的 # 可执行规格(单测 oracle,triton writer 与其逐字节对拍),不可删除。 @@ -1239,7 +1099,6 @@ def pack_indexer_k_to_cache( indexer_k.reshape(-1, self.indexer_head_dim), mem_index.reshape(-1), positions.reshape(-1), - self.full_to_c4_indexs, self.c4_indexer_pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]), self.c4_indexer_pool.page_size, ) diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 164ea7dca0..0cd5401de8 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -85,6 +85,18 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: class DeepseekV4MemOperator(BaseMemManagerOperator): + def copy_mem_to_mem(self, src_mem_index: torch.Tensor, dst_mem_index: torch.Tensor): + """Copy packed history pages; continuation is restored by the request manager.""" + manager = self.mem_manager + src, dst = src_mem_index, dst_mem_index + page_size = manager.page_size + assert src.numel() == dst.numel() and src.numel() % page_size == 0 + src_pages = (src.reshape(-1, page_size)[:, 0] // page_size).to(device="cuda", dtype=torch.int64) + dst_pages = (dst.reshape(-1, page_size)[:, 0] // page_size).to(device="cuda", dtype=torch.int64) + for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + if pool is not None: + pool.buffer.index_copy_(1, dst_pages, pool.buffer.index_select(1, src_pages)) + def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( DeepseekV4MemoryManager, diff --git a/lightllm/common/req_manager/deepseek4.py b/lightllm/common/req_manager/deepseek4.py index e38b8ed9bd..4176db4267 100644 --- a/lightllm/common/req_manager/deepseek4.py +++ b/lightllm/common/req_manager/deepseek4.py @@ -20,8 +20,8 @@ class DeepseekV4PromptCachePayload: """prompt cache 载荷: swa 按页有效性 bitmap 和最后有效页。 - 槽位与 compressor 状态都不进载荷: full_to_swa/full_to_c4/full_to_c128 以 full token 槽位 - 为键(radix 持有 full 槽 ⇒ 映射行存活,free 级联回收);c4 compressor 状态随 swa 页 + 压缩历史由统一 token 页持有;SWA 通过 full slot 映射定位。 + 槽位与 compressor 状态都不进载荷;c4 compressor 状态随 swa 页 生灭。c128 状态按 request ring 寻址;prompt cache 的 256-token 边界同时是 c128 分组边界,命中后新分组会在首次读取前覆写完整 128-token 窗口,因而无需保存状态。 @@ -91,6 +91,7 @@ def __init__( # 出窗回收水位线: -1 表示该 req 尚未见过 prefill chunk(首个 chunk 的 ready_cache_len # 即共享前缀边界,作为永不下探的回收下界)。 self._swa_evict_marks = [-1 for _ in range(max_request_num + 1)] + self._swa_allocated_end = [0 for _ in range(max_request_num + 1)] return # ------------------------------------------------------------------ swa slot prep (per step) @@ -150,6 +151,8 @@ def prepare_prefill_swa( ready_list=ready_list, seq_list=seq_list, ) + for req_idx, seq_len in zip(req_list, seq_list): + self._swa_allocated_end[req_idx] = seq_len return def prepare_decode( @@ -158,16 +161,11 @@ def prepare_decode( b_seq_len_cpu, b_mtp_index_cpu, mem_indexes, - mtp_decode_slot_prepare_indices, - prepare_compress_slots=True, ): """decode 每步槽位 prep。在 BaseModel 的通用 req scatter 与 attention metadata 构建前调用;DeepSeek-V4 MTP draft layer 只需要 SWA 槽位。""" max_mtp_index = int(b_mtp_index_cpu.max().item()) - if mtp_decode_slot_prepare_indices is None: - steps = range(max_mtp_index + 1) - else: - steps = mtp_decode_slot_prepare_indices + steps = range(max_mtp_index + 1) batch_size = b_mtp_index_cpu.shape[0] slots_per_req = max_mtp_index + 1 @@ -184,13 +182,6 @@ def prepare_decode( mem_indexes_by_req[:, step], prev_mem_indexes=mem_indexes_by_req[:, step - 1] if step > 0 else None, ) - if prepare_compress_slots: - self.prepare_decode_compress_slots( - step_req_list, - step_seq_list, - mem_indexes_by_req[:, step], - prev_group_end_mem_indexes=mem_indexes_by_req[:, step - 4] if step >= 4 else None, - ) return def prepare_prefill( @@ -212,12 +203,6 @@ def prepare_prefill( seq_list=seq_list, mem_indexes=mem_indexes, ) - self.prepare_prefill_compress_slots( - req_list=req_list, - ready_list=ready_list, - seq_list=seq_list, - mem_indexes=mem_indexes, - ) return def prepare_pd_decode_cache( @@ -234,12 +219,6 @@ def prepare_pd_decode_cache( assert new_full_slots.numel() == sum(seq_len - ready for ready, seq_len in zip(ready_list, seq_list)) new_full_slots = new_full_slots.reshape(-1).to(self.req_to_token_indexs.device, non_blocking=True) - self.prepare_prefill_compress_slots( - req_list=req_list, - ready_list=ready_list, - seq_list=seq_list, - mem_indexes=new_full_slots, - ) # swa 只保存最后一部分,前面的不需要 swa_start_list = [] @@ -261,8 +240,9 @@ def prepare_pd_decode_cache( ready_list=swa_start_list, seq_list=seq_list, ) - for req_idx, swa_start in zip(req_list, swa_start_list): + for req_idx, swa_start, seq_len in zip(req_list, swa_start_list, seq_list): self._swa_evict_marks[req_idx] = swa_start + self._swa_allocated_end[req_idx] = seq_len return def prepare_decode_swa( @@ -295,6 +275,19 @@ def prepare_decode_swa( self._swa_evict_marks[req_idx] = evict_end if evict_slots: self.mem_manager.evict_swa(torch.cat(evict_slots)) + new_rows = [ + i + for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) + if req_idx != self.HOLD_REQUEST_ID and seq_len > self._swa_allocated_end[req_idx] + ] + if not new_rows: + return + req_list = [req_list[i] for i in new_rows] + seq_list = [seq_list[i] for i in new_rows] + rows = torch.tensor(new_rows, dtype=torch.int64, device=mem_indexes.device) + mem_indexes = mem_indexes[rows] + if prev_mem_indexes is not None: + prev_mem_indexes = prev_mem_indexes[rows] if prev_mem_indexes is None: prev_meta = g_pin_mem_manager.gen_from_list( key="dsv4_swa_decode_prev", @@ -309,6 +302,8 @@ def prepare_decode_swa( mem_indexes, prev_mem_indexes, ) + for req_idx, seq_len in zip(req_list, seq_list): + self._swa_allocated_end[req_idx] = seq_len return def init_compress_state(self, req_idx: int): @@ -321,243 +316,7 @@ def init_compress_state(self, req_idx: int): def finish_cpu_cache_load(self, req_idx: int, loaded_len: int) -> None: """Keep only the final restored 256-token SWA page eligible for radix reuse.""" self._swa_evict_marks[req_idx] = loaded_len - self.get_prompt_cache_page_size() - return - - # ------------------------------------------------------------------ compress slot prep (per step) - def _register_c4_slots(self, full_slots: torch.Tensor, slots: torch.Tensor) -> None: - """写入 full->c4 槽映射并按页累加存活计数。""" - self.mem_manager.full_to_c4_indexs[full_slots] = slots - self.mem_manager.count_c4_slots(slots, 1) - - def _scatter_c4_prefill_slots_batched(self, req_list, ready_list, seq_list, mem_indexes) -> None: - """Batch c4 prefill scatter from the generic preprocess full-slot layout. - - Each group's end token is in the current chunk, so its full slot is addressed directly in - mem_indexes. Only a mid-page continuation reads the previous group's old req-table entry. - New full slots guarantee a fresh mapping; no GPU-to-CPU idempotency check is needed.""" - page = DSV4_C4_PAGE_SIZE - mapping = self.mem_manager.full_to_c4_indexs - device = mapping.device - - plan = [] - mem_offset = 0 - for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): - q_len = seq_len - ready_len - if req_idx == self.HOLD_REQUEST_ID: - mem_offset += q_len - continue - first, last = ready_len // 4, seq_len // 4 - if last <= first: - mem_offset += q_len - continue - plan.append((req_idx, ready_len, mem_offset, first, last)) - mem_offset += q_len - if not plan: - return - - def to_cuda_long(key, data): - return g_pin_mem_manager.gen_from_list(key=key, data=data, dtype=torch.int64).to(device, non_blocking=True) - - reqs, readies, mem_offsets, firsts, lasts = zip(*plan) - counts = [last - first for first, last in zip(firsts, lasts)] - first_pages = [first // page for first in firsts] - page_counts = [((last - 1) // page) - fp + 1 for last, fp in zip(lasts, first_pages)] - page_offsets, total_pages = [], 0 - for n_pages in page_counts: - page_offsets.append(total_pages) - total_pages += n_pages - total_entries = sum(counts) - cont = [(off, req, first) for off, req, first in zip(page_offsets, reqs, firsts) if first % page != 0] - self._realize_c4_pages(total_pages - len(cont)) - - # One pinned H2D copy for all per-request metadata, then per-entry ragged expansion. - meta = to_cuda_long( - "dsv4_c4_prefill_meta", - [x for row in zip(readies, mem_offsets, firsts, first_pages, counts, page_offsets) for x in row], - ).view(-1, 6) - readies_t, mem_offsets_t, firsts_t, first_pages_t, counts_t, page_offsets_t = meta.unbind(1) - seg = torch.repeat_interleave(torch.arange(len(plan), device=device), counts_t, output_size=total_entries) - seg_starts = counts_t.cumsum(0) - counts_t - entries = firsts_t[seg] + torch.arange(total_entries, device=device) - seg_starts[seg] - full_offsets = mem_offsets_t[seg] + entries * 4 + 3 - readies_t[seg] - full_slots = mem_indexes.reshape(-1)[full_offsets] - - # physical base per logical page: fresh pages from one alloc; mid-page continuations read prev - if not cont: - page_bases = self.mem_manager.alloc_c4_pages(total_pages).to(device, non_blocking=True) * page - else: - page_bases = torch.empty(total_pages, dtype=torch.int32, device=device) - new_pos = [ - pos - for off, n_pages, first in zip(page_offsets, page_counts, firsts) - for pos in range(off + (first % page != 0), off + n_pages) - ] - if new_pos: - new_pos_t = to_cuda_long("dsv4_c4_prefill_new_pos", new_pos) - page_bases[new_pos_t] = ( - self.mem_manager.alloc_c4_pages(len(new_pos)).to(device, non_blocking=True) * page - ) - cont_t = to_cuda_long("dsv4_c4_prefill_cont", [x for row in cont for x in row]).view(-1, 3) - prev_slot = mapping[self.req_to_token_indexs[cont_t[:, 1], cont_t[:, 2] * 4 - 1]] - cont_off = ((cont_t[:, 2] - 1) % page).to(torch.int32) - page_bases[cont_t[:, 0]] = prev_slot - cont_off - - page_idx = page_offsets_t[seg] + torch.div(entries, page, rounding_mode="floor") - first_pages_t[seg] - slots = page_bases[page_idx] + (entries % page).to(torch.int32) - self._register_c4_slots(full_slots, slots) - return - - def _scatter_c4_decode_slots( - self, - req_list, - seq_list, - mem_indexes: torch.Tensor, - prev_group_end_mem_indexes: Optional[torch.Tensor] = None, - ) -> None: - page = DSV4_C4_PAGE_SIZE - mapping = self.mem_manager.full_to_c4_indexs - mem_indexes = mem_indexes.reshape(-1) - - cont_rows, cont_prev_pos = [], [] - new_rows = [] - for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): - if req_idx == self.HOLD_REQUEST_ID or seq_len <= 0 or seq_len % 4 != 0: - continue - entry = seq_len // 4 - 1 - offset = entry % page - if offset == 0: - new_rows.append(i) - else: - cont_rows.append(i) - cont_prev_pos.append(entry * 4 - 1) - - if cont_rows: - if prev_group_end_mem_indexes is None: - prev_meta = g_pin_mem_manager.gen_from_list( - key="dsv4_c4_decode_prev", - data=[x for row in zip([req_list[i] for i in cont_rows], cont_prev_pos) for x in row], - dtype=torch.int64, - ).to(mapping.device, non_blocking=True) - prev_meta = prev_meta.view(-1, 2) - prev_full = self.req_to_token_indexs[prev_meta[:, 0], prev_meta[:, 1]] - else: - prev_full = ( - prev_group_end_mem_indexes.reshape(-1) - if len(cont_rows) == len(req_list) - else prev_group_end_mem_indexes.reshape(-1)[cont_rows] - ) - prev_slots = mapping[prev_full] - dst_indexes = mem_indexes if len(cont_rows) == len(req_list) else mem_indexes[cont_rows] - self._register_c4_slots(dst_indexes, prev_slots + 1) - - if new_rows: - self._realize_c4_pages(len(new_rows)) # 兑现: 精确需求, 复用已算的 new_rows - pages = self.mem_manager.alloc_c4_pages(len(new_rows)).to(mapping.device, non_blocking=True) - dst_indexes = mem_indexes if len(new_rows) == len(req_list) else mem_indexes[new_rows] - self._register_c4_slots(dst_indexes, pages * page) - return - - def _scatter_c128_slots(self, full_slots: torch.Tensor) -> None: - """为本批新组末 full 槽分配 c128 槽并写入映射。""" - if full_slots.numel() == 0: - return - full_slots = full_slots.reshape(-1) - self._realize_c128_slots(full_slots.numel()) - new_slots = self.mem_manager.alloc_c128(full_slots.numel()).cuda(non_blocking=True) - self.mem_manager.full_to_c128_indexs[full_slots] = new_slots - return - - def _realize_c4_pages(self, need_pages: int) -> None: - """压缩池兑现 —— 和主池在 prep 里调 free_radix_cache_to_get_enough_token 同一套路: - base_backend admission 已按"空闲+可回收"放行本步请求,这里在真分配前(scatter 已算好 need) - 把可回收的无引用 radix 节点驱逐出来腾出 c4 页,避免 alloc_c4_pages 触底 assert。 - 可回收仍不足时由 admission 的 wait_pause 兜底。""" - if self.mem_manager.n_c4 == 0 or need_pages <= 0: - return - # 延迟 import: infer_batch 在模块顶 import 了 req_manager,顶层 import 会循环引用 - 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_c4_pages(need_pages) - return - - def _realize_c128_slots(self, need_slots: int) -> None: - if self.mem_manager.n_c128 == 0 or need_slots <= 0: - return - 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_c128_slots(need_slots) - return - - def prepare_prefill_compress_slots( - self, - req_list: List[int], - ready_list: List[int], - seq_list: List[int], - mem_indexes: torch.Tensor, - ) -> None: - """prefill prep: 为本 chunk 内的组末 token(位置 (g+1)*ratio-1 ∈ [ready, seq))分配压缩槽, - 组末 full 槽直接从 generic preprocess 的 mem_indexes 取。""" - if self.mem_manager.n_c4 == 0 and self.mem_manager.n_c128 == 0: - return - if self.mem_manager.n_c4 > 0: - self._scatter_c4_prefill_slots_batched(req_list, ready_list, seq_list, mem_indexes) - - if self.mem_manager.n_c128 > 0: - ratio = 128 - full_offsets = [] - mem_offset = 0 - for req_idx, ready_len, seq_len in zip(req_list, ready_list, seq_list): - q_len = seq_len - ready_len - if req_idx == self.HOLD_REQUEST_ID: - mem_offset += q_len - continue - first, last = ready_len // ratio, seq_len // ratio - if last > first: - full_offsets.extend( - mem_offset + (entry + 1) * ratio - 1 - ready_len for entry in range(first, last) - ) - mem_offset += q_len - if full_offsets: - offsets = g_pin_mem_manager.gen_from_list( - key="dsv4_c128_prefill_offsets", data=full_offsets, dtype=torch.int64 - ).to(mem_indexes.device, non_blocking=True) - self._scatter_c128_slots(mem_indexes.reshape(-1)[offsets]) - return - - def prepare_decode_compress_slots( - self, - req_list: List[int], - seq_list: List[int], - mem_indexes: torch.Tensor, - prev_group_end_mem_indexes: Optional[torch.Tensor] = None, - ) -> None: - """decode prep: 本步 token 关闭一个组(seq_len % ratio == 0)时为其分配压缩槽并 scatter。 - 组末 full 槽即本步的 mem_index。 - 从 CPU 镜像读 seq_len/req_idx(host 算术,无 D2H);非关组步 rows 为空 => 不调 _scatter,零同步。""" - if self.mem_manager.n_c4 == 0 and self.mem_manager.n_c128 == 0: - return - if self.mem_manager.n_c4 > 0: - self._scatter_c4_decode_slots( - req_list, - seq_list, - mem_indexes, - prev_group_end_mem_indexes=prev_group_end_mem_indexes, - ) - - if self.mem_manager.n_c128 > 0: - ratio = 128 - rows = [ - i - for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) - if req_idx != self.HOLD_REQUEST_ID and seq_len > 0 and seq_len % ratio == 0 - ] - if rows: - full_slots = mem_indexes.reshape(-1) - if len(rows) != len(req_list): - full_slots = full_slots[rows] - self._scatter_c128_slots(full_slots) + self._swa_allocated_end[req_idx] = loaded_len return def alloc(self): @@ -569,6 +328,7 @@ def alloc(self): def clear_runtime_state(self, req_idx: int): # swa 槽位本身由 mem_manager.free 级联回收(随 full 槽位),这里只复位出窗水位线。 self._swa_evict_marks[req_idx] = -1 + self._swa_allocated_end[req_idx] = 0 return def get_prompt_cache_value_ops(self): @@ -661,4 +421,5 @@ def free_req(self, free_req_index: int): def free_all(self): super().free_all() self._swa_evict_marks = [-1 for _ in range(self.max_request_num + 1)] + self._swa_allocated_end = [0 for _ in range(self.max_request_num + 1)] return diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index 05e573f7ce..9f4a13c5e2 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -113,7 +113,6 @@ def init_some_extra_state(self, model): self.dsv4_sparse_req_idx, self.position_ids, self.req_manager.req_to_token_indexs, - self.mem_manager.full_to_c128_indexs, 128, self.dsv4_c128_indices, self.dsv4_c128_lengths, diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index 5b675ec882..e2555a5283 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -319,9 +319,14 @@ def fused_compress( out_buffer: torch.Tensor = None, ): mem_manager = infer_state.mem_manager + out_slots = torch.where( + (infer_state.position_ids + 1) % compress_ratio == 0, + infer_state.mem_index.reshape(-1) // compress_ratio, + -1, + ) if is_in_indexer: assert compress_ratio == 4, "只有 c4(CSA) 层有 indexer-K" - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx) state_ring = mem_manager.c4_state_ring if out_buffer is None: @@ -333,12 +338,12 @@ def fused_compress( out_page_size = 1 else: if compress_ratio == 4: - out_slots = mem_manager.full_to_c4_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c4_state_buffer(layer_idx) state_ring = mem_manager.c4_state_ring out_page_size = mem_manager.c4_pool.page_size elif compress_ratio == 128: - out_slots = mem_manager.full_to_c128_indexs[infer_state.mem_index.reshape(-1)] + state_buffer = mem_manager.get_c128_state_buffer(layer_idx) state_ring = mem_manager.c128_state_ring out_page_size = mem_manager.c128_pool.page_size diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 08a07b10c4..579bebbfea 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -905,7 +905,6 @@ def _c4_indices(self, infer_state: DeepseekV4InferStateInfo, idx_q_fp8, weights, infer_state.dsv4_sparse_req_idx, positions, infer_state.req_manager.req_to_token_indexs, - mem_manager.full_to_c4_indexs, 4, slots, lengths, diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 6e8828a946..4e58b55b7b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -118,7 +118,7 @@ def _init_mem_manager(self): mem_fraction=self.mem_fraction, memory_reservations=reservations, ) - self.req_manager.mem_manager = self.mem_manager + self.req_manager.bind_mem_manager(self.mem_manager) return def _get_post_profile_memory_reservations(self): @@ -293,59 +293,31 @@ def prepare_mtp_layer_hidden(self, layer_index: int, hidden): streams = hidden.view(-1, self.config["hc_mult"], self.config["hidden_size"]) return streams.mean(dim=1) - def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: - """Commit DSV4 derived slots before BaseModel pads or scatters the generic input.""" - if model_input.batch_size == 0: + def _prepare_dsv4_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor) -> None: + if model_input.batch_size == 0 or (model_input.is_prefill and self.is_mtp_draft_model): return - if model_input.is_prefill and self.is_mtp_draft_model: - return - if model_input.mem_indexes is None: - model_input.mem_indexes = model_input.mem_indexes_cpu.cuda(non_blocking=True) - + # Runtime inputs retain CPU mirrors. Synthetic warmup inputs are created on GPU. + req_ids = model_input.b_req_idx_cpu + seq_lens = model_input.b_seq_len_cpu + if req_ids is None: + req_ids = model_input.b_req_idx.cpu() + seq_lens = model_input.b_seq_len.cpu() if model_input.is_prefill: - if model_input.mem_indexes_cpu is None: - model_input.b_req_idx_cpu = model_input.b_req_idx.detach().cpu() - model_input.b_seq_len_cpu = model_input.b_seq_len.detach().cpu() - self.req_manager.prepare_prefill( - b_req_idx_cpu=model_input.b_req_idx_cpu, - b_ready_cache_len_cpu=model_input.b_ready_cache_len, - b_seq_len_cpu=model_input.b_seq_len_cpu, - mem_indexes=model_input.mem_indexes, - ) - return - - if model_input.mtp_decode_slot_prepare_indices == (): - return - # GPU-only decode inputs are CUDA Graph warmup/HOLD layouts. Runtime - # DSV4 draft inputs carry CPU mirrors from the proposer and prepare slots here. - if model_input.mem_indexes_cpu is None: - return - self.req_manager.prepare_decode( - model_input.b_req_idx_cpu, - model_input.b_seq_len_cpu, - model_input.b_mtp_index_cpu, - model_input.mem_indexes, - model_input.mtp_decode_slot_prepare_indices, - prepare_compress_slots=not self.is_mtp_draft_model, - ) - return - - @torch.no_grad() - def forward(self, model_input: ModelInput): - self._prepare_dsv4_slots(model_input) - return super().forward(model_input) - - @torch.no_grad() - def microbatch_overlap_prefill(self, model_input0: ModelInput, model_input1: ModelInput): - self._prepare_dsv4_slots(model_input0) - self._prepare_dsv4_slots(model_input1) - return super().microbatch_overlap_prefill(model_input0, model_input1) - - @torch.no_grad() - def microbatch_overlap_decode(self, model_input0: ModelInput, model_input1: ModelInput): - self._prepare_dsv4_slots(model_input0) - self._prepare_dsv4_slots(model_input1) - return super().microbatch_overlap_decode(model_input0, model_input1) + ready_lens = model_input.b_ready_cache_len_cpu + if ready_lens is None: + ready_lens = model_input.b_ready_cache_len.cpu() + token_num = int((seq_lens - ready_lens).sum()) + self.req_manager.prepare_prefill(req_ids, ready_lens, seq_lens, mem_indexes[:token_num]) + else: + mtp_indices = model_input.b_mtp_index_cpu + if mtp_indices is None: + mtp_indices = model_input.b_mtp_index.cpu() + self.req_manager.prepare_decode(req_ids, seq_lens, mtp_indices, mem_indexes[: len(req_ids)]) + + def _select_mem_indexes(self, model_input: ModelInput): + mem_indexes = super()._select_mem_indexes(model_input) + self._prepare_dsv4_slots(model_input, mem_indexes) + return mem_indexes def _init_to_get_rotary(self): # Interleaved (GPT-J) rope. Build complex64 freqs_cis tables (_freqs_cis_*) following the diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py index dae5121ab9..3d9202e292 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py @@ -9,7 +9,6 @@ def _build_compress_index_kernel( pos_ptr, req_to_token_ptr, req_to_token_stride0, - full_to_c_ptr, index_ptr, index_stride0, length_ptr, @@ -30,7 +29,7 @@ def _build_compress_index_kernel( end_pos = e * RATIO + (RATIO - 1) safe_pos = tl.where(valid, end_pos, 0) full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_pos, mask=valid, other=0).to(tl.int64) - c_slot = tl.load(full_to_c_ptr + full_slot, mask=valid, other=-1).to(tl.int32) + c_slot = tl.where(valid, full_slot // RATIO, -1).to(tl.int32) tl.store(index_ptr + t * index_stride0 + e, c_slot, mask=e_mask) if eb == 0: @@ -41,20 +40,14 @@ def build_compress_index( req_idx: torch.Tensor, positions: torch.Tensor, req_to_token_indexs: torch.Tensor, - full_to_c_indexs: torch.Tensor, ratio: int, index: torch.Tensor, length: torch.Tensor, ): - """Fused two-level group-end gather for the c4/c128 compressed-entry index tables. + """Derive compressed entries from group-end full slots in 256-token pages. - For token t (at request `req_idx[t]`, absolute `positions[t]`) and compressed entry e: - slot[t, e] = full_to_c[ req_to_token[req, e*ratio + (ratio-1)] ] (the group-end token's full slot) - with slot = -1 where e >= (pos+1)//ratio (beyond the causal compressed length) or where the - full->c map is unset. Writes index [T, cap] and length [T] = clamp((pos+1)//ratio, 1). - - Replaces the eager _gather_compress_slots/_c128/c4-causal torch chain. The caller owns the - output storage, so this wrapper does not allocate on the hot path. + Only entries closed at the logical query position are visible. Reserved + capacity beyond that position is masked to -1, regardless of its contents. """ T = positions.shape[0] cap = index.shape[1] @@ -67,7 +60,6 @@ def build_compress_index( positions, req_to_token_indexs, req_to_token_indexs.stride(0), - full_to_c_indexs, index, index.stride(0), length, diff --git a/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py index 33b015947e..0159d39d51 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py @@ -9,14 +9,12 @@ @triton.jit def _pool_pages_kernel( full_slots, - full_to_pool, pool, pool_stride0, pool_stride1, paired_pool, paired_pool_stride0, paired_pool_stride1, - full_to_c128, c128_pool, c128_pool_stride0, c128_pool_stride1, @@ -74,7 +72,7 @@ def _pool_pages_kernel( + first_full_offset + gpu_page_i64 * full_offset_per_gpu_page ).to(tl.int64) - pool_slot = tl.load(full_to_pool + full_slot).to(tl.int64) + pool_slot = full_slot // 4 physical_page = pool_slot // pool_page_size mask = offsets < gpu_page_nbytes pool_ptr = pool + layer_i64 * pool_stride0 + physical_page * pool_stride1 + offsets_i64 @@ -128,7 +126,7 @@ def _pool_pages_kernel( + c128_first_full_offset + c128_row_i64 * c128_full_offset_per_row ).to(tl.int64) - c128_pool_slot = tl.load(full_to_c128 + c128_full_slot).to(tl.int64) + c128_pool_slot = c128_full_slot // 128 c128_physical_page = c128_pool_slot // c128_pool_page_size c128_token_in_page = c128_pool_slot % c128_pool_page_size c128_pool_page = c128_pool + layer_i64 * c128_pool_stride0 + c128_physical_page * c128_pool_stride1 @@ -179,7 +177,6 @@ def copy_pool_pages( mode, *, full_slots, - mapping, pool, staging, page_num, @@ -196,7 +193,6 @@ def copy_pool_pages( paired_pool=None, paired_section_offset=0, paired_section_layer_nbytes=0, - c128_mapping=None, c128_pool=None, c128_row_num=0, c128_first_full_offset=0, @@ -224,14 +220,12 @@ def copy_pool_pages( ) _pool_pages_kernel[(page_num * grid_layer_num, grid_gpu_page_num, byte_blocks_per_gpu_page)]( full_slots, - mapping, pool, pool.stride(0) if has_pool else 0, pool.stride(1) if has_pool else 0, paired_pool, paired_pool.stride(0) if paired_pool is not None else 0, paired_pool.stride(1) if paired_pool is not None else 0, - c128_mapping, c128_buffer, c128_buffer.stride(0) if has_c128 else 0, c128_buffer.stride(1) if has_c128 else 0, diff --git a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py index 724c1afed7..6955a35643 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py @@ -36,14 +36,12 @@ def _scatter_staging_to_cpu_kernel( @triton.jit def _pack_gpu_cache_to_staging_kernel( full_slots, - full_to_c4, c4_pool, c4_pool_stride0, c4_pool_stride1, c4_indexer_pool, c4_indexer_pool_stride0, c4_indexer_pool_stride1, - full_to_c128, c128_pool, c128_pool_stride0, c128_pool_stride1, @@ -119,7 +117,7 @@ def _pack_gpu_cache_to_staging_kernel( full_slot = tl.load( full_slots + logical_page_i64 * token_page_size + 3 + gpu_page_i64 * history_block_size ).to(tl.int64) - pool_slot = tl.load(full_to_c4 + full_slot).to(tl.int64) + pool_slot = full_slot // 4 physical_page = pool_slot // c4_pool_page_size offsets = byte_block * BLOCK + offsets_base offsets_i64 = offsets.to(tl.int64) @@ -163,7 +161,7 @@ def _pack_gpu_cache_to_staging_kernel( full_slot = tl.load( full_slots + logical_page_i64 * token_page_size + c128_ratio - 1 + row_i64 * c128_ratio ).to(tl.int64) - pool_slot = tl.load(full_to_c128 + full_slot).to(tl.int64) + pool_slot = full_slot // 128 physical_page = pool_slot // c128_pool_page_size token_in_page = pool_slot % c128_pool_page_size offsets_i64 = offsets_base.to(tl.int64) @@ -515,14 +513,12 @@ def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, sta _pack_gpu_cache_to_staging_kernel[(page_num * programs_per_page,)]( full_slots, - mem_manager.full_to_c4_indexs if has_c4 else None, c4_pool, c4_pool.stride(0) if has_c4 else 0, c4_pool.stride(1) if has_c4 else 0, c4_indexer_pool, c4_indexer_pool.stride(0) if has_c4 else 0, c4_indexer_pool.stride(1) if has_c4 else 0, - mem_manager.full_to_c128_indexs if has_c128 else None, c128_pool, c128_pool.stride(0) if has_c128 else 0, c128_pool.stride(1) if has_c128 else 0, diff --git a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py index 25724000bc..baa99bf684 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py @@ -9,7 +9,6 @@ def _fwd_kernel_destindex_copy_indexer_k_dsv4( K, Mem_index, Positions, - Full_to_c4, O_fp8, O_f32, stride_k_bs, @@ -28,9 +27,7 @@ def _fwd_kernel_destindex_copy_indexer_k_dsv4( return full_slot = tl.load(Mem_index + cur_index).to(tl.int64) - dest_index = tl.load(Full_to_c4 + full_slot).to(tl.int64) - if dest_index < 0: - return + dest_index = full_slot // COMPRESS_RATIO page = dest_index // PAGE_SIZE token_in_page = dest_index % PAGE_SIZE @@ -54,7 +51,6 @@ def destindex_copy_indexer_k_dsv4( K: torch.Tensor, MemIndex: torch.Tensor, Positions: torch.Tensor, - FullToC4: torch.Tensor, O_buffer: torch.Tensor, page_size: int, ): @@ -63,8 +59,7 @@ def destindex_copy_indexer_k_dsv4( K: [T, 128] bf16 unquantized indexer keys. MemIndex: [T] int — full-token slots for the current rows. Positions: [T] int — logical token positions; only c4 group-end rows are written. - FullToC4: [full_pool_size + 1] int — full-token slot to c4-pool slot mapping. - Negative mappings are skipped. + C4 slots are group-end full-token slots divided by four. O_buffer: [num_pages, bytes_per_page] uint8 — one layer's slab from the c4 indexer PackedPagePool (128B fp8 data region + 4B fp32 scale tail per token). @@ -89,7 +84,6 @@ def destindex_copy_indexer_k_dsv4( K, MemIndex, Positions, - FullToC4, flat.view(torch.float8_e4m3fn), flat.view(torch.float32), K.stride(0), diff --git a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py index 4b929235b8..5df62f4131 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py @@ -8,8 +8,8 @@ _C4_RATIO = 4 _C128_RATIO = 128 _BYTE_BLOCK = 8192 -# Per-source row: c4 map/data/indexer, c128 map/data, SWA map/data, c4 state/indexer state. -_SOURCE_POOL_PTR_COUNT = 9 +# Per-source row: c4 data/indexer, c128 data, SWA map/data, c4 state/indexer state. +_SOURCE_POOL_PTR_COUNT = 7 # Per-task row: source manager index, token count, source full-slot pointer, destination full-slot pointer. _TASK_META_WIDTH = 4 @@ -19,14 +19,12 @@ def _copy_dsv4_dp_caches_kernel( source_pool_ptrs, task_meta, history_meta, - dst_full_to_c4, dst_c4_pool, dst_c4_pool_stride0, dst_c4_pool_stride1, dst_c4_indexer_pool, dst_c4_indexer_pool_stride0, dst_c4_indexer_pool_stride1, - dst_full_to_c128, dst_c128_pool, dst_c128_pool_stride0, dst_c128_pool_stride1, @@ -88,14 +86,13 @@ def _copy_dsv4_dp_caches_kernel( if HAS_C4: if layer < c4_layer_num: - src_full_to_c4 = tl.load(source_ptr_row).to(tl.pointer_type(tl.int32)) - src_c4_pool = tl.load(source_ptr_row + 1).to(tl.pointer_type(tl.uint8)) - src_c4_indexer_pool = tl.load(source_ptr_row + 2).to(tl.pointer_type(tl.uint8)) + src_c4_pool = tl.load(source_ptr_row + 0).to(tl.pointer_type(tl.uint8)) + src_c4_indexer_pool = tl.load(source_ptr_row + 1).to(tl.pointer_type(tl.uint8)) full_offset = history_block * history_block_size + c4_ratio - 1 src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_pool_slot = tl.load(src_full_to_c4 + src_full_slot).to(tl.int64) - dst_pool_slot = tl.load(dst_full_to_c4 + dst_full_slot).to(tl.int64) + src_pool_slot = src_full_slot // c4_ratio + dst_pool_slot = dst_full_slot // c4_ratio src_page = src_pool_slot // c4_pool_page_size dst_page = dst_pool_slot // c4_pool_page_size @@ -129,14 +126,13 @@ def _copy_dsv4_dp_caches_kernel( if HAS_C128: if layer < c128_layer_num: - src_full_to_c128 = tl.load(source_ptr_row + 3).to(tl.pointer_type(tl.int32)) - src_c128_pool = tl.load(source_ptr_row + 4).to(tl.pointer_type(tl.uint8)) + src_c128_pool = tl.load(source_ptr_row + 2).to(tl.pointer_type(tl.uint8)) for row in tl.static_range(0, 2): full_offset = history_block * history_block_size + (row + 1) * c128_ratio - 1 src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_pool_slot = tl.load(src_full_to_c128 + src_full_slot).to(tl.int64) - dst_pool_slot = tl.load(dst_full_to_c128 + dst_full_slot).to(tl.int64) + src_pool_slot = src_full_slot // c128_ratio + dst_pool_slot = dst_full_slot // c128_ratio src_page = src_pool_slot // c128_pool_page_size dst_page = dst_pool_slot // c128_pool_page_size src_token = src_pool_slot % c128_pool_page_size @@ -169,8 +165,8 @@ def _copy_dsv4_dp_caches_kernel( src_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) dst_full_slots = tl.load(task_row + 3).to(tl.pointer_type(tl.int32)) source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count - src_full_to_swa = tl.load(source_ptr_row + 5).to(tl.pointer_type(tl.int32)) - src_swa_pool = tl.load(source_ptr_row + 6).to(tl.pointer_type(tl.uint8)) + src_full_to_swa = tl.load(source_ptr_row + 3).to(tl.pointer_type(tl.int32)) + src_swa_pool = tl.load(source_ptr_row + 4).to(tl.pointer_type(tl.uint8)) if task_pid < swa_program_num: page = task_pid % 2 @@ -199,8 +195,8 @@ def _copy_dsv4_dp_caches_kernel( layer = state_pid // 4 row_i64 = row.to(tl.int64) layer_i64 = layer.to(tl.int64) - src_c4_state = tl.load(source_ptr_row + 7).to(tl.pointer_type(tl.uint8)) - src_c4_indexer_state = tl.load(source_ptr_row + 8).to(tl.pointer_type(tl.uint8)) + src_c4_state = tl.load(source_ptr_row + 5).to(tl.pointer_type(tl.uint8)) + src_c4_indexer_state = tl.load(source_ptr_row + 6).to(tl.pointer_type(tl.uint8)) full_offset = token_num - 4 + row_i64 src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) @@ -274,14 +270,12 @@ def copy_dsv4_dp_caches( source_pool_ptrs, task_meta, history_meta, - dst_mem_manager.full_to_c4_indexs if has_c4 else None, dst_c4_pool, dst_c4_pool.stride(0) if has_c4 else 0, dst_c4_pool.stride(1) if has_c4 else 0, dst_c4_indexer_pool, dst_c4_indexer_pool.stride(0) if has_c4 else 0, dst_c4_indexer_pool.stride(1) if has_c4 else 0, - dst_mem_manager.full_to_c128_indexs if has_c128 else None, dst_c128_pool, dst_c128_pool.stride(0) if has_c128 else 0, dst_c128_pool.stride(1) if has_c128 else 0, diff --git a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py index 54d2958836..7bf6910eef 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py @@ -9,7 +9,6 @@ def _build_c4_indexer_page_table_kernel( c4_len_ptr, # [batch] int req_to_token_ptr, req_to_token_stride0, - full_to_c4_ptr, page_table_ptr, # [batch, page_cap] int32 page_cap, hold_req_id, @@ -29,7 +28,7 @@ def _build_c4_indexer_page_table_kernel( mask=active, other=0, ).to(tl.int64) - c4_slot0 = tl.load(full_to_c4_ptr + full_slot0, mask=active, other=0).to(tl.int64) + c4_slot0 = full_slot0 // RATIO phys_page = c4_slot0 // PAGE_SIZE tl.store(page_table_ptr + r * page_cap + p, tl.where(active, phys_page, 0).to(tl.int32)) @@ -66,7 +65,6 @@ def build_c4_indexer_page_table( c4_len, req_to_token_indexs, req_to_token_indexs.stride(0), - mem_manager.full_to_c4_indexs, page_table, page_cap, int(hold_req_id), diff --git a/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py index b26a3d81b2..33e5eb5e67 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py @@ -322,7 +322,6 @@ def _copy_pd_cache_page( copy_pool_pages( mode, full_slots=full_slots, - mapping=mem_manager.full_to_c4_indexs if c4_work else None, pool=c4_pool.buffer if c4_work else None, staging=staging, page_num=1, @@ -338,7 +337,6 @@ def _copy_pd_cache_page( paired_pool=c4_indexer_pool.buffer if c4_work else None, paired_section_offset=layout.c4_indexer_offset, paired_section_layer_nbytes=layout.c4_indexer_layer_nbytes, - c128_mapping=mem_manager.full_to_c128_indexs if c128_work else None, c128_pool=c128_pool if c128_work else None, c128_row_num=c128_row_num, c128_first_full_offset=_C128_RATIO - 1, diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py index 81aa1a269b..2739dbd20a 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py @@ -1,5 +1,7 @@ import torch +from lightllm.common.basemodel.batch_objs import PostLayerOutput + from lightllm.models.deepseek_v4.layer_infer.hyper_connection import hc_head, hc_post from lightllm.models.deepseek_v4_dspark.layer_weights.pre_and_post_layer_weight import ( DeepseekV4DSparkPreAndPostLayerWeight, @@ -36,7 +38,7 @@ def token_forward( layer_weight: DeepseekV4DSparkPreAndPostLayerWeight, ): if infer_state.is_prefill: - return input_embdings.new_empty((0,)) + return PostLayerOutput(logits=input_embdings.new_empty((0,))) if isinstance(input_embdings, tuple): streams = hc_post(*input_embdings) @@ -79,4 +81,4 @@ def token_forward( confidence_logits=confidence_logits, ) # The proposer consumes token ids directly; keep only the row dimension for graph unpadding. - return local_logits.new_empty((token_num, 1)) + return PostLayerOutput(logits=local_logits.new_empty((token_num, 1))) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 3e0afe271b..7ac58c364d 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -183,25 +183,21 @@ def _init_custom(self): layer.sin_compress_table = self._sin_cached_compress self.layers_infer[0].context_wkv_weight = self.context_wkv_weight - def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: - if not model_input.is_prefill: - if model_input.mtp_draft_input_hiddens is not None: - # Target verify already prepared these shared full/SWA slots. - # This pass only commits target hiddens into the draft layers. - model_input.mtp_decode_slot_prepare_indices = () - elif model_input.mem_indexes_cpu is not None: - # Proposal-owned full slots live only until the next verify. - # Put each request's complete block in a private scratch page; - # this needs neither accepted-length D2H nor host seq metadata. - ( - model_input.mtp_draft_swa_pages_cpu, - model_input.mtp_draft_swa_pages, - ) = self.mem_manager.alloc_dspark_swa_block( - mem_indexes=model_input.mem_indexes, - block_size=self.block_size, - ) - model_input.mtp_decode_slot_prepare_indices = () - return super()._prepare_dsv4_slots(model_input) + def _prepare_dsv4_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor) -> None: + # Target-hidden commits use target SWA; proposal blocks own separate scratch pages. + return + + @torch.no_grad() + def forward(self, model_input: ModelInput): + if model_input.is_prefill or model_input.mtp_draft_input_hiddens is not None: + return super().forward(model_input) + pages_cpu, pages = self.mem_manager.alloc_dspark_swa_block(model_input.batch_size, self.block_size) + model_input.mtp_draft_swa_pages_cpu = pages_cpu + model_input.mtp_draft_swa_pages = pages + try: + return super().forward(model_input) + finally: + self.mem_manager.free_dspark_swa_block(pages_cpu) def _decode(self, model_input: ModelInput) -> ModelOutput: if model_input.mtp_draft_input_hiddens is None: diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index e232f01d86..c324798e8b 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -217,6 +217,8 @@ def _launch_subprocesses(args: StartArgs): if args.page_size < 1: raise ValueError(f"--page_size must be >= 1, got {args.page_size}") + if get_model_type(args.model_dir) == "deepseek_v4" and args.page_size != 256: + raise ValueError("DeepSeek-V4 requires --page_size 256") if args.run_mode in ("prefill", "decode"): assert args.pd_kv_page_size % args.page_size == 0, "--pd_kv_page_size must be divisible by --page_size" diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index 62a971ab7a..40caf8ee10 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -779,44 +779,6 @@ def release_mem(mem_index): self.mem_manager.free(mem_index) return - def _free_radix_full_nodes_until(self, allocator, need: int) -> None: - """DeepSeek-V4 压缩池(c4/c128)兑现: 沿 LRU 序逐个驱逐 ref_count==0 的整个 full radix 节点, - 经 mem_manager.free() 级联回收其 c4 页 / c128 槽(evict_c4/evict_c128),每驱逐一个就复查 - *真实* allocator(不靠计数,稳),直到够或已无可驱逐的无引用节点。后者(空闲+可回收仍不足) - 由上游 base_backend admission 的 wait_pause 兜底,allocator 的 assert 是最后防线。""" - if self.mem_manager is None or allocator is None: - return - while allocator.can_use_mem_size < need: - # 无可驱逐的无引用 token => 停(admission 应已 wait_pause) - if self.tree_total_tokens_num.arr[0] <= self.refed_tokens_num.arr[0]: - # 兜底没兜住:admission/realize 估算漂移了。打日志便于定位(否则只会撞下游隐晦的 - # allocator "error alloc state" assert)。 - logger.warning( - f"dsv4 compress-pool realize could not free enough: need={need} " - f"free={allocator.can_use_mem_size} tree_total={self.tree_total_tokens_num.arr[0]} " - f"refed={self.refed_tokens_num.arr[0]} (admission should have paused this req)" - ) - return - release_mems = [] - # 复用已测的 evict():弹一个 LRU、ref==0 的叶子(>=1 token),其 full 槽经 free 级联回收压缩槽 - self.evict(1, lambda mem_index: release_mems.append(mem_index)) - self.mem_manager.free(torch.concat(release_mems)) - return - - def free_radix_cache_to_get_enough_c4_pages(self, need_pages: int) -> None: - allocator = getattr(self.mem_manager, "c4_page_allocator", None) if self.mem_manager is not None else None - if allocator is None or need_pages <= 0: - return - self._free_radix_full_nodes_until(allocator, need_pages) - return - - def free_radix_cache_to_get_enough_c128_slots(self, need_slots: int) -> None: - allocator = getattr(self.mem_manager, "c128_allocator", None) if self.mem_manager is not None else None - if allocator is None or need_slots <= 0: - return - self._free_radix_full_nodes_until(allocator, need_slots) - return - class _RadixCacheReadOnlyClient: """ diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 962436f216..72b563699b 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -414,13 +414,7 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo if paused_reqs: # pause_reqs 可能刚刚释放了 KV,因此恢复前重新读取实时可用容量。 can_alloc_token_num = self.get_can_alloc_token_num() - can_alloc_dsv4_swa_page_num = can_alloc_dsv4_c4_page_num = can_alloc_dsv4_c128_slot_num = None - if self.is_deepseek_v4: - ( - can_alloc_dsv4_swa_page_num, - can_alloc_dsv4_c4_page_num, - can_alloc_dsv4_c128_slot_num, - ) = self.get_can_alloc_dsv4_page_and_slot_num() + can_alloc_dsv4_swa_page_num = self.get_can_alloc_dsv4_swa_page_num() if self.is_deepseek_v4 else None for req in paused_reqs: # 暂停恢复保持原有的保守语义:只有当前完整序列所需的 KV 页面都有足够空间时才恢复, @@ -429,18 +423,9 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo if alloc_token_num > can_alloc_token_num: break - swa_page_num = c4_page_num = c128_slot_num = 0 - if ( - can_alloc_dsv4_swa_page_num is not None - or can_alloc_dsv4_c4_page_num is not None - or can_alloc_dsv4_c128_slot_num is not None - ): - swa_page_num, c4_page_num, c128_slot_num = req.get_dsv4_recover_need_page_and_slot_num() - if can_alloc_dsv4_swa_page_num is not None and swa_page_num > can_alloc_dsv4_swa_page_num: - break - if can_alloc_dsv4_c4_page_num is not None and c4_page_num > can_alloc_dsv4_c4_page_num: - break - if can_alloc_dsv4_c128_slot_num is not None and c128_slot_num > can_alloc_dsv4_c128_slot_num: + if self.is_deepseek_v4: + swa_page_num = req.get_dsv4_recover_need_swa_page_num() + if swa_page_num > can_alloc_dsv4_swa_page_num: break if g_infer_context.is_hybrid_att_model: @@ -456,10 +441,6 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo can_alloc_token_num -= alloc_token_num if can_alloc_dsv4_swa_page_num is not None: can_alloc_dsv4_swa_page_num -= swa_page_num - if can_alloc_dsv4_c4_page_num is not None: - can_alloc_dsv4_c4_page_num -= c4_page_num - if can_alloc_dsv4_c128_slot_num is not None: - can_alloc_dsv4_c128_slot_num -= c128_slot_num return def get_can_alloc_token_num(self): @@ -470,28 +451,11 @@ def get_can_alloc_token_num(self): ) return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num - def get_can_alloc_dsv4_page_and_slot_num(self): - self.req_manager: DeepseekV4ReqManager - mem_manager = self.req_manager.mem_manager - radix_cache_unref_page_num = 0 - radix_cache_unref_token_num = 0 + def get_can_alloc_dsv4_swa_page_num(self): + pages = int(self.req_manager.mem_manager.swa_page_allocator.can_use_mem_size) if self.radix_cache is not None: - radix_cache_unref_page_num = self.radix_cache.get_unrefed_swa_pages_num() - radix_cache_unref_token_num = ( - self.radix_cache.get_tree_total_tokens_num() - self.radix_cache.get_refed_tokens_num() - ) - swa_page_num = int(mem_manager.swa_page_allocator.can_use_mem_size) + radix_cache_unref_page_num - - c4_page_num = 0 - if mem_manager.c4_page_allocator is not None: - c4_page_num = int(mem_manager.c4_page_allocator.can_use_mem_size) + int( - radix_cache_unref_token_num // self.req_manager.get_prompt_cache_page_size() - ) - - c128_slot_num = 0 - if mem_manager.c128_allocator is not None: - c128_slot_num = int(mem_manager.c128_allocator.can_use_mem_size) + int(radix_cache_unref_token_num // 128) - return swa_page_num, c4_page_num, c128_slot_num + pages += self.radix_cache.get_unrefed_swa_pages_num() + return pages def save_hybrid_state_to_cache(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """Snapshot request-level attention state at big/small-page boundaries.""" @@ -689,10 +653,6 @@ def __init__( if g_infer_context.is_deepseek_v4: mem_manager = g_infer_context.req_manager.mem_manager self.dsv4_swa_page_size: int = mem_manager.swa_pool.page_size - self.dsv4_c4_page_size: int = ( - mem_manager.c4_pool.page_size if mem_manager.c4_page_allocator is not None else 0 - ) - self.dsv4_has_c128: bool = mem_manager.c128_allocator is not None self._init_all_state() @@ -1094,31 +1054,22 @@ def _kv_cache_alloc_need(self, target_kv_len: int) -> int: assert alloc_token_num % page_size == 0 return alloc_token_num - def get_dsv4_prefill_need_page_and_slot_num(self, is_chuncked_prefill: bool) -> Tuple[int, int, int]: + def get_dsv4_prefill_need_swa_page_num(self, is_chuncked_prefill: bool) -> int: start = self.cur_kv_len end = self.get_chuncked_input_token_len() if is_chuncked_prefill else self.get_cur_total_len() if end <= start: - return 0, 0, 0 + return 0 first_new_page = (start + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size last_page = (end - 1) // self.dsv4_swa_page_size swa_page_num = last_page - first_new_page + 1 - c4_page_num = 0 - first, last = start // 4, end // 4 - if last > first: - # Safe upper bound: touched c4 pages, including a possible already-allocated continuation page. - c4_page_num = (last - 1) // self.dsv4_c4_page_size - first // self.dsv4_c4_page_size + 1 - - c128_slot_num = max(0, end // 128 - start // 128) if self.dsv4_has_c128 else 0 - return swa_page_num, c4_page_num, c128_slot_num + return swa_page_num - def get_dsv4_recover_need_page_and_slot_num(self) -> Tuple[int, int, int]: - swa_page_num, c4_page_num, c128_slot_num = self.get_dsv4_prefill_need_page_and_slot_num( - is_chuncked_prefill=False - ) + def get_dsv4_recover_need_swa_page_num(self) -> int: + swa_page_num = self.get_dsv4_prefill_need_swa_page_num(is_chuncked_prefill=False) if swa_page_num == 0 or self.args.disable_chunked_prefill: - return swa_page_num, c4_page_num, c128_slot_num + return swa_page_num # C4/C128 accumulate across recovery chunks; only SWA is evicted chunk by chunk. req_manager: DeepseekV4ReqManager = g_infer_context.req_manager @@ -1132,28 +1083,19 @@ def get_dsv4_recover_need_page_and_slot_num(self) -> Tuple[int, int, int]: max_prefill_token_num + int(req_manager.sliding_window) + 2 * prompt_cache_page_size, ) swa_page_num = (peak_token_num + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size - return swa_page_num, c4_page_num, c128_slot_num + return swa_page_num - def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: + def get_dsv4_decode_need_swa_page_num(self) -> int: seq_len = self.get_cur_total_len() if seq_len <= 0: - return 0, 0, 0 + return 0 swa_page_num = 0 - c4_page_num = 0 - c128_slot_num = 0 # Main model prepares current token plus draft-verify rows: SWA + compressed slots. for step in range(self.mtp_step + 1): cur_seq_len = seq_len + step if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: swa_page_num += 1 - if cur_seq_len % 4 == 0: - entry = cur_seq_len // 4 - 1 - if entry % self.dsv4_c4_page_size == 0: - c4_page_num += 1 - if self.dsv4_has_c128 and cur_seq_len % 128 == 0: - c128_slot_num += 1 - if self.args.mtp_mode == "dspark": # DSpark proposal allocates one private SWA scratch page per request. swa_page_num += 1 @@ -1164,7 +1106,7 @@ def get_dsv4_decode_need_page_and_slot_num(self) -> Tuple[int, int, int]: cur_seq_len = seq_len + step if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: swa_page_num += 1 - return swa_page_num, c4_page_num, c128_slot_num + return swa_page_num class InferReqUpdatePack: diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index e25dfa1f85..edf16809ea 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -752,14 +752,8 @@ def _get_classed_reqs( can_alloc_token_num = g_infer_context.get_can_alloc_token_num() is_deepseek_v4 = self.is_deepseek_v4 can_alloc_dsv4_swa_page_num = None - can_alloc_dsv4_c4_page_num = None - can_alloc_dsv4_c128_slot_num = None if is_deepseek_v4: - ( - can_alloc_dsv4_swa_page_num, - can_alloc_dsv4_c4_page_num, - can_alloc_dsv4_c128_slot_num, - ) = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() + can_alloc_dsv4_swa_page_num = g_infer_context.get_can_alloc_dsv4_swa_page_num() for req_obj in ready_reqs: @@ -795,20 +789,14 @@ def _get_classed_reqs( _, alloc_token_num = req_obj.decode_need_token_num() can_run = alloc_token_num <= can_alloc_token_num if can_run and is_deepseek_v4: - swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_decode_need_page_and_slot_num() - can_run = ( - swa_page_num <= can_alloc_dsv4_swa_page_num - and c4_page_num <= can_alloc_dsv4_c4_page_num - and c128_slot_num <= can_alloc_dsv4_c128_slot_num - ) + swa_page_num = req_obj.get_dsv4_decode_need_swa_page_num() + can_run = swa_page_num <= can_alloc_dsv4_swa_page_num if can_run: self._alloc_req_kv_mem(req_obj, alloc_token_num, no_blcoking_copy=True) decode_reqs.append(req_obj) can_alloc_token_num -= alloc_token_num if is_deepseek_v4: can_alloc_dsv4_swa_page_num -= swa_page_num - can_alloc_dsv4_c4_page_num -= c4_page_num - can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: if self.args.run_mode == "decode": @@ -842,14 +830,8 @@ def _get_classed_reqs( continue can_run = alloc_token_num <= can_alloc_token_num if can_run and is_deepseek_v4: - swa_page_num, c4_page_num, c128_slot_num = req_obj.get_dsv4_prefill_need_page_and_slot_num( - is_chuncked_prefill=is_chuncked_prefill - ) - can_run = ( - swa_page_num <= can_alloc_dsv4_swa_page_num - and c4_page_num <= can_alloc_dsv4_c4_page_num - and c128_slot_num <= can_alloc_dsv4_c128_slot_num - ) + swa_page_num = req_obj.get_dsv4_prefill_need_swa_page_num(is_chuncked_prefill=is_chuncked_prefill) + can_run = swa_page_num <= can_alloc_dsv4_swa_page_num if can_run: self._alloc_req_kv_mem(req_obj, alloc_token_num, no_blcoking_copy=True) prefill_tokens += token_num @@ -857,8 +839,6 @@ def _get_classed_reqs( can_alloc_token_num -= alloc_token_num if is_deepseek_v4: can_alloc_dsv4_swa_page_num -= swa_page_num - can_alloc_dsv4_c4_page_num -= c4_page_num - can_alloc_dsv4_c128_slot_num -= c128_slot_num else: if wait_pause_count < pause_max_req_num: req_obj.wait_pause = True diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index f9af71b97b..f865a637a2 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -48,10 +48,8 @@ def init_dsv4_cache_transfer(self, mem_managers: List[MemoryManager]) -> None: has_c128 = mem_manager.c128_pool is not None pointer_rows.append( [ - mem_manager.full_to_c4_indexs.data_ptr() if has_c4 else 0, mem_manager.c4_pool.buffer.data_ptr() if has_c4 else 0, mem_manager.c4_indexer_pool.buffer.data_ptr() if has_c4 else 0, - mem_manager.full_to_c128_indexs.data_ptr() if has_c128 else 0, mem_manager.c128_pool.buffer.data_ptr() if has_c128 else 0, mem_manager.full_to_swa_indexs.data_ptr(), mem_manager.swa_pool.buffer.data_ptr(), @@ -90,11 +88,7 @@ def build_shared_kv_trans_tasks( trans_tasks: List[TransTask] = [] if self.backend.is_deepseek_v4: dsv4_mem_manager = self.backend.model.mem_manager - ( - dsv4_swa_capacity, - dsv4_c4_capacity, - dsv4_c128_capacity, - ) = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() + dsv4_swa_capacity = g_infer_context.get_can_alloc_dsv4_swa_page_num() dsv4_prompt_page_size = self.backend.model.req_manager.get_prompt_cache_page_size() rank_max_radix_cache_lens = np.max( @@ -112,13 +106,7 @@ def build_shared_kv_trans_tasks( can_alloc_dsv4_cache = True if self.backend.is_deepseek_v4 and trans_size > 0: need_swa_pages = dsv4_prompt_page_size // dsv4_mem_manager.swa_pool.page_size - need_c4_pages = trans_size // dsv4_prompt_page_size if dsv4_mem_manager.c4_pool is not None else 0 - need_c128_slots = trans_size // 128 if dsv4_mem_manager.c128_pool is not None else 0 - can_alloc_dsv4_cache = ( - dsv4_swa_capacity >= need_swa_pages - and dsv4_c4_capacity >= need_c4_pages - and dsv4_c128_capacity >= need_c128_slots - ) + can_alloc_dsv4_cache = dsv4_swa_capacity >= need_swa_pages target_kv_len = req.cur_kv_len + trans_size alloc_token_num = req._kv_cache_alloc_need(target_kv_len) if trans_size > 0 else 0 @@ -137,8 +125,6 @@ def build_shared_kv_trans_tasks( mem_indexes = mem_indexes[:trans_size] if self.backend.is_deepseek_v4: dsv4_swa_capacity -= need_swa_pages - dsv4_c4_capacity -= need_c4_pages - dsv4_c128_capacity -= need_c128_slots max_kv_len_dp_rank = self.shared_req_infos.arr[req_index, :, self._KV_LEN_INDEX].argmax() max_kv_len_req_idx = int(self.shared_req_infos.arr[req_index, max_kv_len_dp_rank, self._REQ_IDX_INDEX]) max_kv_len_mem_manager_index = max_kv_len_dp_rank * self.backend.dp_world_size + self.backend.rank_in_dp diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py index f36d72af07..5c1746474b 100644 --- a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py @@ -281,27 +281,21 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): layout = mem_manager.cpu_cache_layout requested_end = int(page_len_list[len(page_list) - 1]) if requested_end > gpu_kv_len: - swa_capacity, c4_capacity, c128_capacity = g_infer_context.get_can_alloc_dsv4_page_and_slot_num() + swa_capacity = g_infer_context.get_can_alloc_dsv4_swa_page_num() loadable_end = mem_manager.get_loadable_cpu_cache_end( gpu_kv_len, requested_end, idle_token_num, swa_capacity, - c4_capacity, - c128_capacity, ) loadable_end = self._get_image_safe_load_end(req, gpu_kv_len, loadable_end, layout.token_page_size) if loadable_end != 0: token_num = loadable_end - gpu_kv_len full_need = token_num swa_need = 2 - c4_need = token_num // 256 if mem_manager.n_c4 else 0 - c128_need = token_num // 128 if mem_manager.n_c128 else 0 if self.backend.radix_cache is not None: radix_cache = self.backend.radix_cache radix_cache.free_radix_cache_to_get_enough_token(full_need) - radix_cache.free_radix_cache_to_get_enough_c4_pages(c4_need) - radix_cache.free_radix_cache_to_get_enough_c128_slots(c128_need) swa_shortage = swa_need - int(mem_manager.swa_page_allocator.can_use_mem_size) if swa_shortage > 0: radix_cache.free_unreferenced_swa_pages(swa_shortage) @@ -311,8 +305,6 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): loadable_end, int(mem_manager.allocator.can_use_mem_size), int(mem_manager.swa_page_allocator.can_use_mem_size), - int(mem_manager.c4_page_allocator.can_use_mem_size) if mem_manager.n_c4 else 0, - int(mem_manager.c128_allocator.can_use_mem_size) if mem_manager.n_c128 else 0, ) loadable_end = self._get_image_safe_load_end( req, gpu_kv_len, loadable_end, layout.token_page_size @@ -336,6 +328,7 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): ] = plan.mem_indexes self.backend.model.req_manager.finish_cpu_cache_load(req.req_idx, loaded_end) req.cur_kv_len = loaded_end + req.hold_kv_len = loaded_end idle_token_num -= token_num if is_master_in_dp: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index 4374b76cf4..f404cf66b2 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -176,23 +176,11 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: swa_page = mem_manager.swa_pool.page_size swa_need = max(0, (input_len - 1) // swa_page - (swa_start + swa_page - 1) // swa_page + 1) - _, c4_need, c128_need = req_obj.get_dsv4_prefill_need_page_and_slot_num(is_chuncked_prefill=False) - c4_allocator = mem_manager.c4_page_allocator - c128_allocator = mem_manager.c128_allocator - - # D ingress 会立即分配派生槽,必须在任何分配前先兑现并检查实际容量。 if self.radix_cache is not None: - self.radix_cache.free_radix_cache_to_get_enough_c4_pages(c4_need) - self.radix_cache.free_radix_cache_to_get_enough_c128_slots(c128_need) swa_shortage = swa_need - mem_manager.swa_page_allocator.can_use_mem_size if swa_shortage > 0: self.radix_cache.free_unreferenced_swa_pages(swa_shortage) - - if ( - swa_need > mem_manager.swa_page_allocator.can_use_mem_size - or (c4_allocator is not None and c4_need > c4_allocator.can_use_mem_size) - or (c128_allocator is not None and c128_need > c128_allocator.can_use_mem_size) - ): + if swa_need > mem_manager.swa_page_allocator.can_use_mem_size: return False mem_indexes = self._alloc_req_kv_mem(req_obj, need_mem_size) diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index b4912b2d12..db774f62aa 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -112,7 +112,6 @@ def propose_next_overlap( ): model_input.input_ids = token_ids model_input.mtp_draft_input_hiddens = model_output.mtp_collector.spec_hidden - model_input.mtp_decode_slot_prepare_indices = () proposal_token_ids = target_next_token_ids0.new_empty((req_num, draft_step)) schedule_scores = ( @@ -204,8 +203,6 @@ def propose_next_overlap( ) 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] - model_input.mem_indexes_cpu = None - model_input.mtp_decode_slot_prepare_indices = None if self.backend.is_deepseek_v4 and req_num > 0: accepted_tail_rows_cpu = accepted_tail_rows_cpu_by_batch[batch_index] model_input.b_req_idx_cpu = model_input.b_req_idx_cpu.index_select(0, accepted_tail_rows_cpu) diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index e279816670..1f554ee3f6 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -75,7 +75,6 @@ def propose_next( 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 - verify_draft_input.mtp_decode_slot_prepare_indices = () extend_output = draft_model.forward(verify_draft_input) # 只在 req_num 行 logits 上进行 argmax,避免为未接受的 verify 行执行 diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index 2c9dbd7644..4f68933be5 100644 --- a/test/benchmark/static_inference/static_benchmark.py +++ b/test/benchmark/static_inference/static_benchmark.py @@ -432,7 +432,9 @@ def _run_mtp_draft_decode( model_output: ModelOutput, step_width: int, ): - draft_input = model_input.make_mtp_draft_input() + draft_input = copy.copy(model_input) + draft_input.b_seq_len = model_input.b_seq_len.clone() + draft_input.b_seq_len_cpu = model_input.b_seq_len_cpu.clone() draft_output = model_output draft_next_ids = self._argmax_ids(model_output.logits).cuda(non_blocking=True) generated = [draft_next_ids.detach()] @@ -447,6 +449,7 @@ def _run_mtp_draft_decode( if self.args.mtp_mode.startswith("eagle") and step + 1 < self.args.mtp_step: draft_input.b_seq_len += 1 + draft_input.b_seq_len_cpu += 1 draft_input.max_kv_seq_len += 1 return torch.stack(generated[:step_width], dim=1) @@ -526,12 +529,6 @@ def _materialize_cached_prefix_extra_slots( seq_list=seq_list, mem_indexes=mem_indexes_gpu[:, swa_ready_len:].contiguous(), ) - req_manager.prepare_prefill_compress_slots( - req_list=req_list, - ready_list=[0] * batch_size, - seq_list=seq_list, - mem_indexes=mem_indexes_gpu, - ) def _cached_prefix_swa_ready_len(self, cached_len: int) -> int: req_manager: DeepseekV4ReqManager = self.model.req_manager diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index fc260d93bc..206438d986 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -15,15 +15,13 @@ def test_dspark_decode_admission_reserves_one_swa_scratch_page(): req.args = SimpleNamespace(mtp_mode="eagle") req.mtp_step = 1 req.dsv4_swa_page_size = 128 - req.dsv4_c4_page_size = 64 - req.dsv4_has_c128 = True req.get_cur_total_len = lambda: 10 - normal_need = req.get_dsv4_decode_need_page_and_slot_num() + normal_need = req.get_dsv4_decode_need_swa_page_num() req.args.mtp_mode = "dspark" - dspark_need = req.get_dsv4_decode_need_page_and_slot_num() + dspark_need = req.get_dsv4_decode_need_swa_page_num() - assert dspark_need == (normal_need[0] + 1, normal_need[1], normal_need[2]) + assert dspark_need == normal_need + 1 @pytest.mark.parametrize("mtp_step", [1, 4, 5]) @@ -107,7 +105,6 @@ def test_dspark_cuda_graph_padding_extends_only_gpu_scratch_pages(): b_position_delta=torch.zeros(batch_size, dtype=torch.int32), b_shared_seq_len=torch.zeros(batch_size, dtype=torch.int32), b_shared_radix_node_id=torch.full((batch_size,), -1, dtype=torch.int64), - mem_indexes=torch.arange(batch_size, dtype=torch.int32), is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], mtp_draft_swa_pages_cpu=pages_cpu, @@ -143,7 +140,6 @@ def test_dspark_empty_decode_padding_builds_one_hold_block(): 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=[], ) @@ -156,48 +152,37 @@ def test_dspark_empty_decode_padding_builds_one_hold_block(): assert padded_input.input_ids.tolist() == [1] * block_size assert padded_input.b_req_idx.tolist() == [127] * block_size assert padded_input.b_seq_len.tolist() == [2] * block_size - assert padded_input.mem_indexes.tolist() == [255] * block_size -def test_dspark_scratch_cleanup_keeps_page_ids_on_cpu(): - 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, - ) +@pytest.mark.parametrize("fails", [False, True]) +def test_dspark_forward_releases_only_scratch_pages(monkeypatch, fails): + from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel - class FakeMemManager: - def __init__(self): - self.scratch_frees = [] - self.normal_frees = [] - - def free_dspark_swa_block(self, mem_indexes_cpu, pages_cpu): - self.scratch_frees.append((mem_indexes_cpu.clone(), pages_cpu.clone())) - - def free(self, mem_indexes_cpu): - self.normal_frees.append(mem_indexes_cpu.clone()) - - mem_manager = FakeMemManager() - backend = SimpleNamespace(model=SimpleNamespace(req_manager=SimpleNamespace(mem_manager=mem_manager))) - scratch_full = torch.tensor([10, 11, 12, 13, 14], dtype=torch.int32) - scratch_pages = torch.tensor([7], dtype=torch.int32) - normal_full = torch.tensor([20, 21], dtype=torch.int32) - - mtp_utils.free_mem_indexes( - backend=backend, - extra_mem_indexes_cpu=[ - MtpMemIndexesToFree( - mem_indexes_cpu=scratch_full, - swa_pages_cpu=scratch_pages, - ), - MtpMemIndexesToFree(mem_indexes_cpu=normal_full), - ], + pages_cpu = torch.tensor([7], dtype=torch.int32) + frees = [] + model = DeepseekV4DSparkModel.__new__(DeepseekV4DSparkModel) + model.block_size = 5 + model.mem_manager = SimpleNamespace( + alloc_dspark_swa_block=lambda count, width: (pages_cpu, pages_cpu.clone()), + free_dspark_swa_block=lambda pages: frees.append(pages), ) + model_input = SimpleNamespace(is_prefill=False, mtp_draft_input_hiddens=None, batch_size=5) + + def forward(self, inputs): + assert inputs.mtp_draft_swa_pages_cpu is pages_cpu + if fails: + raise RuntimeError("draft failed") + return "output" - assert len(mem_manager.scratch_frees) == 1 - torch.testing.assert_close(mem_manager.scratch_frees[0][0], scratch_full) - torch.testing.assert_close(mem_manager.scratch_frees[0][1], scratch_pages) - assert len(mem_manager.normal_frees) == 1 - torch.testing.assert_close(mem_manager.normal_frees[0], normal_full) + monkeypatch.setattr(DeepseekV4TpPartModel, "forward", forward) + if fails: + with pytest.raises(RuntimeError, match="draft failed"): + model.forward(model_input) + else: + assert model.forward(model_input) == "output" + assert len(frees) == 1 + assert frees[0] is pages_cpu @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -279,7 +264,7 @@ def test_dspark_swa_block_uses_one_scratch_page_per_request(): ) mem_indexes = torch.tensor([3, 4, 5, 8, 9, 10], dtype=torch.int64, device="cuda") - pages_cpu, pages = manager.alloc_dspark_swa_block(mem_indexes=mem_indexes, block_size=3) + pages_cpu, pages = manager.alloc_dspark_swa_block(token_num=mem_indexes.numel(), block_size=3) torch.testing.assert_close( manager.full_to_swa_indexs[mem_indexes], diff --git a/unit_tests/common/test_deepseek4_paged_cache.py b/unit_tests/common/test_deepseek4_paged_cache.py new file mode 100644 index 0000000000..bdb0b6b0f4 --- /dev/null +++ b/unit_tests/common/test_deepseek4_paged_cache.py @@ -0,0 +1,192 @@ +import json +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager +from lightllm.common.req_manager import DeepseekV4ReqManager +from lightllm.models.deepseek_v4.triton_kernel.build_compress_index_dsv4 import build_compress_index +from lightllm.utils.envs_utils import get_env_start_args + + +@pytest.fixture +def cache(monkeypatch, tmp_path): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for the packed cache kernels") + (tmp_path / "config.json").write_text(json.dumps({"vocab_size": 100})) + monkeypatch.setenv("LIGHTLLM_CURRENT_RANK_IN_NODE", "0") + monkeypatch.setenv("LIGHTLLM_UNIQUE_SERVICE_NAME_ID", "test_dsv4_paged_cache") + monkeypatch.setenv( + "LIGHTLLM_START_ARGS", + json.dumps( + { + "page_size": 256, + "model_dir": str(tmp_path), + "penalty_counter_mode": "cpu_counter", + "mtp_step": 3, + "mtp_dynamic_verify": False, + "enable_ep_moe": False, + } + ), + ) + get_env_start_args.cache_clear() + manager = DeepseekV4MemoryManager( + 8192, + torch.bfloat16, + 1, + 512, + 3, + compress_rates=[4, 128, 0], + max_request_num=2, + mtp_step=3, + swa_full_tokens_ratio=1.0, + ) + requests = DeepseekV4ReqManager(2, 4096, manager, sliding_window=128) + yield manager, requests + torch.cuda.synchronize() + get_env_start_args.cache_clear() + + +def _reserve(manager, requests, length): + req_idx = requests.alloc() + held = (length + 255) // 256 * 256 + slots = manager.alloc(held).cuda() + requests.req_to_token_indexs[req_idx, :held] = slots + return req_idx, slots + + +@pytest.mark.parametrize("ratio", [4, 128]) +def test_noncontiguous_pages_obey_logical_compression_boundaries(cache, ratio): + manager, requests = cache + req_idx, allocated = _reserve(manager, requests, 768) + slots = allocated.view(3, 256)[torch.tensor([2, 0, 1], device="cuda")].reshape(-1) + requests.req_to_token_indexs[req_idx, :768] = slots + positions = torch.tensor([0, 2, 3, 126, 127, 254, 255, 256, 511, 512], device="cuda") + indexes = torch.empty((len(positions), 768 // ratio), dtype=torch.int32, device="cuda") + lengths = torch.empty_like(positions, dtype=torch.int32) + build_compress_index( + torch.full_like(positions, req_idx), positions, requests.req_to_token_indexs, ratio, indexes, lengths + ) + group_slots = slots[ratio - 1 :: ratio] // ratio + for row, position in enumerate(positions.tolist()): + count = (position + 1) // ratio + torch.testing.assert_close(indexes[row, :count], group_slots[:count]) + assert indexes[row, count:].eq(-1).all() + assert lengths[row].item() == max(1, count) + manager.free(allocated) + assert manager.allocator.can_use_mem_size == manager.size + + +def test_indexer_writer_only_writes_closed_groups(cache): + manager, requests = cache + _, slots = _reserve(manager, requests, 512) + slots = slots.view(2, 256).flip(0).reshape(-1) + keys = torch.randn((259, 128), dtype=torch.bfloat16, device="cuda") + positions = torch.arange(259, dtype=torch.int32, device="cuda") + manager.c4_indexer_pool.buffer.zero_() + manager.pack_indexer_k_to_cache(0, slots[:259], positions, keys) + closed = torch.arange(3, 259, 4, device="cuda") + actual = manager.c4_indexer_pool.read(0, slots[closed] // 4) + expected = manager._pack_indexer_k(keys[closed]) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + # Positions 256..258 cannot publish their reserved next compressed slot. + assert manager.c4_indexer_pool.read(0, slots[256:257] // 4).eq(0).all() + + +def test_speculative_retry_reuses_swa_without_freeing_token_pages(cache): + manager, requests = cache + req_idx, slots = _reserve(manager, requests, 512) + + def cpu(data): + return torch.tensor(data, dtype=torch.int32) + + requests.prepare_prefill(cpu([req_idx]), cpu([0]), cpu([254]), slots[:254]) + for sequences in ([255, 256, 257, 258], [255, 256, 257, 258], [256, 257, 258, 259]): + indexes = slots[cpu(sequences).cuda().long() - 1] + requests.prepare_decode(cpu([req_idx] * 4), cpu(sequences), cpu([0, 1, 2, 3]), indexes) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages - 3 + assert manager.swa_page_live_count.sum().item() == max(sequences) + assert manager.allocator.can_use_mem_size == manager.size - 512 + requests.free([req_idx], slots) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + assert manager.swa_page_live_count.eq(0).all() + + +def test_prefill_chunk_keeps_all_new_swa_rows(cache): + manager, requests = cache + req_idx, slots = _reserve(manager, requests, 2048) + requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048]), slots) + assert manager.full_to_swa_indexs[slots].ge(0).all() + assert manager.swa_page_live_count.sum().item() == 2048 + requests.free([req_idx], slots) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +def test_history_copy_preserves_packed_data_and_scale_regions(cache): + manager, requests = cache + _, src = _reserve(manager, requests, 512) + _, dst = _reserve(manager, requests, 512) + src = src.view(2, 256).flip(0).reshape(-1) + for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + pool.buffer.random_(0, 256) + manager.operator.copy_mem_to_mem(src_mem_index=src, dst_mem_index=dst) + for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + torch.testing.assert_close( + pool.buffer[:, src[::256].long() // 256], pool.buffer[:, dst[::256].long() // 256], rtol=0, atol=0 + ) + assert manager.full_to_swa_indexs[dst].eq(-1).all() + + +def test_hold_page_and_dspark_scratch_have_separate_capacity(cache): + manager, requests = cache + hold = requests.req_to_token_indexs[requests.HOLD_REQUEST_ID] + torch.testing.assert_close( + hold[:256], torch.arange(manager.size, manager.size + 256, dtype=torch.int32, device="cuda") + ) + assert manager.full_to_swa_indexs[hold].ge(manager.swa_size).all() + assert (hold // 4).lt(manager.c4_pool.num_pages * 64).all() + assert (hold // 128).lt(manager.c128_pool.num_pages * 2).all() + free_tokens = manager.allocator.can_use_mem_size + for _ in range(3): + pages_cpu, pages = manager.alloc_dspark_swa_block(16, 8) + assert pages.numel() == 2 + assert manager.allocator.can_use_mem_size == free_tokens + manager.free_dspark_swa_block(pages_cpu) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +def test_cpu_cache_roundtrip_uses_derived_history_slots(cache): + manager, requests = cache + req_idx, src = _reserve(manager, requests, 2048) + requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048]), src) + for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool, manager.swa_pool): + pool.buffer.random_(0, 256) + manager.c4_state_buffer.uniform_() + manager.c4_indexer_state_buffer.uniform_() + staging = torch.empty((1, manager.cpu_cache_layout.page_nbytes), dtype=torch.uint8, device="cuda") + manager.operator.pack_cpu_cache_pages(src.view(1, -1), staging) + cpu_page = torch.empty(staging.shape, dtype=torch.uint8, pin_memory=True) + cpu_page.copy_(staging) + plan = manager.prepare_cpu_cache_load(2048, 2048) + manager.operator.load_cpu_cache_pages( + plan, torch.tensor([0], dtype=torch.int32, device="cuda"), SimpleNamespace(cpu_kv_cache_tensor=cpu_page) + ) + manager.commit_cpu_cache_load_plan(plan) + for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): + torch.testing.assert_close( + pool.read(0, src[ratio - 1 :: ratio].long() // ratio), + pool.read(0, plan.mem_indexes[ratio - 1 :: ratio].long() // ratio), + rtol=0, + atol=0, + ) + src_swa = manager.full_to_swa_indexs[src[-256:]].long() + dst_swa = manager.full_to_swa_indexs[plan.mem_indexes[-256:]].long() + for layer in range(manager.layer_num): + torch.testing.assert_close( + manager.swa_pool.read(layer, src_swa), manager.swa_pool.read(layer, dst_swa), rtol=0, atol=0 + ) + manager.free(torch.cat([src, plan.mem_indexes])) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages From 38847af46f9d540a0334587eec1e0ef3d8ac16ae Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 19:18:28 +0800 Subject: [PATCH 192/214] Integrate DeepSeek-V4 private continuation with hybrid checkpoints --- docs/kv_cache_page_size.md | 26 + .../deepseek4_mem_manager.py | 298 ++-------- .../kv_cache_mem_manager/operator/deepseek.py | 14 +- .../operator/linear_att.py | 1 + lightllm/common/req_manager/__init__.py | 2 +- lightllm/common/req_manager/deepseek4.py | 516 +++++------------ lightllm/common/req_manager/hybrid_base.py | 20 +- lightllm/common/req_manager/linear_att.py | 12 +- .../common/state_cache_manager/__init__.py | 7 +- .../common/state_cache_manager/deepseek4.py | 30 + lightllm/models/deepseek_v4/infer_struct.py | 5 +- .../deepseek_v4/layer_infer/compressor.py | 29 +- .../layer_infer/transformer_layer_infer.py | 5 +- lightllm/models/deepseek_v4/model.py | 11 +- .../triton_kernel/build_dspark_swa_index.py | 28 +- .../triton_kernel/build_swa_index_dsv4.py | 34 +- .../deepseek_v4/triton_kernel/cpu_cache_io.py | 34 +- .../deepseek_v4/triton_kernel/dp_cache_io.py | 47 +- .../deepseek_v4/triton_kernel/pd_cache_io.py | 60 +- .../models/deepseek_v4_dspark/infer_struct.py | 4 +- lightllm/models/deepseek_v4_dspark/model.py | 2 +- lightllm/models/deepseek_v4_mtp/model.py | 6 - lightllm/server/api_start.py | 22 +- lightllm/server/core/objs/req.py | 3 +- .../router/dynamic_prompt/radix_cache.py | 301 +--------- .../server/router/model_infer/infer_batch.py | 185 +++---- .../model_infer/mode_backend/base_backend.py | 18 +- .../dp_backend/dp_shared_kv_trans.py | 97 +++- .../mode_backend/dsv4_multi_level_kv_cache.py | 64 ++- .../pd/decode_node_impl/decode_impl.py | 50 +- .../pd/prefill_node_impl/prefill_impl.py | 2 +- lightllm/utils/config_utils.py | 2 +- lightllm/utils/kv_cache_utils.py | 17 +- test/unit/test_deepseek_v4_dspark.py | 50 +- .../common/test_deepseek4_paged_cache.py | 518 +++++++++++++++++- .../deepseek_v4/test_vision_integration.py | 101 +++- .../test_paged_side_allocations.py | 2 + unit_tests/server/test_pd_start_args.py | 43 ++ 38 files changed, 1318 insertions(+), 1348 deletions(-) create mode 100644 lightllm/common/state_cache_manager/deepseek4.py diff --git a/docs/kv_cache_page_size.md b/docs/kv_cache_page_size.md index 76e99a88d7..4910463014 100644 --- a/docs/kv_cache_page_size.md +++ b/docs/kv_cache_page_size.md @@ -43,3 +43,29 @@ diverse mode 仍只支持 `page_size=1`,启动阶段会对其他取值明确 `[HOLD, HOLD+1, ..., HOLD+page_size-1]` 循环填充,而不是重复同一个 token 索引。 - Radix 子节点用“首个完整 token 页”作为键,避免不同序列仅首 token 相同造成页级分支冲突。 - 非法的 `page_size < 1` 以及尚未支持的功能组合在模型加载前失败。 + +## DeepSeek-V4 + +V4 首版整合固定使用 256-token 分配页和 256-token 小页。启动时需要添加: + +```bash +--page_size 256 --linear_att_hash_page_size 256 +``` + +启用每 2048 token 一个大页 checkpoint 时,再添加 `--linear_att_page_block_num 8`。 +省略该参数时沿用主线默认值,关闭大页 checkpoint。PD Decode 节点仍按主线规则关闭大页。 +启用大页和 CPU cache 时,`--cpu_cache_token_page_size` 必须等于大页间隔;未指定时自动设置。 + +- 统一 token 页持有压缩历史。闭合分组的 C4/C128 槽分别由组末 full slot 除以 4/128 得到; + 不再维护独立压缩槽映射或 allocator。预留容量不改变 attention 的逻辑可见长度。 +- packed 格式保持独立:SWA 每页 128 槽,C4 每页 64 槽,C128 每页 2 槽。 + 历史复制、CPU cache 和 PD 使用 V4 专用算子,不能直接按普通 KV 的 token 维度复制。 +- SWA 页和 compressor 运行态归请求所有,保留完整 prefill chunk 所需的数据。 + 两个请求命中同一前缀时,共享压缩历史,分别恢复自己的 SWA 和 compressor 状态。 +- 大小页 CPU checkpoint 只保存末尾 256 token 的 SWA、C4 和 indexer continuation, + 不重复保存压缩历史。256 边界已闭合 C128 分组,下一组覆盖状态后再读取。 +- PD 支持任意传输终点,额外保存未闭合分组以及最后可缓存边界的 continuation。 + PD continuation 布局有变化,P/D 节点需要使用相同代码版本。 + +验证覆盖分页边界、packed history 复制、大小页命中和分叉、推测拒绝后重试、CUDA Graph 索引、 +CPU/PD pack-unpack、跨进程 NCCL checkpoint 传输及资源回收。完整模型精度和性能仍需另行实测。 diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 5d38884ea9..ea3aae2a5f 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -2,7 +2,7 @@ import torch from dataclasses import dataclass -from typing import List, Optional, Sequence, Union +from typing import List, Optional, Sequence from .mem_manager import MemoryManager from .operator import DeepseekV4MemOperator from .allocator import KvCacheAllocator @@ -177,6 +177,26 @@ class DeepseekV4CpuCacheLayout(_DeepseekV4CacheLayout): page_nbytes: int + @classmethod + def load_from_args(cls): + from lightllm.utils.config_utils import ( + get_config_json, + get_layer_num, + get_head_dim, + get_deepseek_v4_compress_rates, + ) + from lightllm.utils.envs_utils import get_added_mtp_kv_layer_num + + args = get_env_start_args() + config = get_config_json(args.model_dir) + layer_num = get_layer_num(args.model_dir) + get_added_mtp_kv_layer_num() + return cls.from_compress_rates( + get_deepseek_v4_compress_rates(config, layer_num), + token_page_size=args.cpu_cache_token_page_size, + head_dim=get_head_dim(args.model_dir), + indexer_head_dim=config["index_head_dim"], + ) + @classmethod def from_compress_rates( cls, @@ -252,7 +272,8 @@ def from_compress_rates( swa_layer_nbytes = swa_gpu_pages_per_page * swa_gpu_page_nbytes swa_nbytes = history["layer_num"] * swa_layer_nbytes - c4_state_rows = DSV4_C4_STATE_RING - 1 + # Four rows at the aligned checkpoint, plus up to seven live-tail rows. + c4_state_rows = 4 + DSV4_C4_STATE_RING - 1 c4_state_row_nbytes = 4 * head_dim * torch._utils._element_size(torch.float32) c4_state_offset = history["swa_offset"] + swa_nbytes c4_state_nbytes = history["n_c4"] * c4_state_rows * c4_state_row_nbytes @@ -297,11 +318,7 @@ class DeepseekV4CpuCacheLoadPlan: history_full_slots: torch.Tensor history_c4_slots: Optional[torch.Tensor] history_c128_slots: Optional[torch.Tensor] - resume_full_slots: torch.Tensor resume_swa_slots: torch.Tensor - resume_full_slots_long: torch.Tensor - resume_swa_pages: torch.Tensor - resume_swa_page_deltas: torch.Tensor class PackedPagePool: @@ -377,17 +394,13 @@ class DeepseekV4MemoryManager(MemoryManager): 与兄弟 manager 一致的 token-slot 设计;req 索引的表都在 DeepseekV4ReqManager。 - - ``swa_pool``: 584B packed latent,所有层。池子小于 full token 空间;prep 阶段 - ``alloc_swa_prefill/decode`` 按**页**(128 槽,位置对齐: slot(p)=page_base+p%128)分配, - 映射记录到 ``full_to_swa_indexs``(以 full token 槽位为键)。出窗槽位由 DeepseekV4ReqManager - 在 prep 阶段批量惰性回收(``evict_swa``,页存活计数减到 0 才整页归还);full 槽位释放时 - ``free`` 级联回收对应 swa 槽,所以 radix 驱逐/请求释放/暂停无需任何额外协议。 - 页 allocator 触底时先走 swa free hook(radix 对 ref==0 节点 free)再 assert。 - 没有 ring buffer,prefill chunk 大小不受 sliding_window 限制。 + - ``swa_pool``: packed latent for all layers. The request manager owns its + physical pages and request-private page table. Prefill retains the entire + current chunk, without a fixed-size ring that could overwrite unread KV. - ``c4_pool``/``c128_pool``: 压缩 latent,按 qwen3next 的层号压实手法只为压缩层建层; c4 另带 packed indexer-K 池。统一 256-token 页拥有三个 packed 历史页; 压缩槽位直接由组末 full slot 除以压缩率得到,随 full 页分配和释放。 - - 写入走标准 operator 路径(``pack_mla_kv_to_cache``),内部为 triton packed writer; + - 写入走模型专用的 fused norm/RoPE packed writer,显式接收本轮 SWA 槽; torch codecs 保留为 ABI 的可执行规格(单测 oracle)。 """ @@ -535,17 +548,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): self.swa_page_allocator = KvCacheAllocator( self.swa_num_pages, shared_name=f"{server}_dsv4_swa_can_use_page_num_{rank_in_node}" ) - # 页存活计数 = 指向该页的有效 full_to_swa 行数;减到 0 归还 allocator(出窗逐 token - # 回收下,「部分出窗页」计数 > 0 自然受保护)。下标含 HOLD 页(只读不增减)。 - self.swa_page_live_count = torch.zeros((self.swa_pool.num_pages,), dtype=torch.int32, device="cuda") - # swa free hook(可选): 页 allocator 触底时回调(radix 对 ref==0 节点 free swa 页), - # 由 backend 在 radix cache 创建后 register;assert 仍是最后防线。 - self._free_radix_unreferenced_swa_fn = None self.HOLD_TOKEN_MEMINDEX = size - self.full_to_swa_indexs = torch.full((size + self.page_size,), -1, dtype=torch.int32, device="cuda") - self.full_to_swa_indexs[size:] = ( - self.swa_size + torch.arange(self.page_size, device="cuda") % DSV4_SWA_PAGE_SIZE - ) self.c4_size = _ceil_div(size, 4) self.c128_size = _ceil_div(size, 128) @@ -572,7 +575,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): scale_bytes=DSV4_INDEXER_SCALE_BYTES, ) # c4 compressor 在途状态(attention + indexer): swa 页派生寻址(翻译③),随 swa 页 - # 生灭 -> radix 命中零拷贝续算。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), + # 生灭;radix 命中从 CPU checkpoint 恢复。行数 = 页数*ring + ring(HOLD 页) + 1(哨兵), # 取整到 ratio;末行哨兵 kv=0/score=-inf(KVAndScore.clear 语义),其余行由内核在 # 组起点覆写,无需按页清零。last_dim = 2*coff*head_dim(overlap coff=2)。 state_rows = self._paged_state_rows(self.swa_num_pages, self.c4_state_ring, 4) @@ -608,6 +611,15 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): assert self.c4_indexer_pool.bytes_per_page == layout.c4_indexer_gpu_page_nbytes assert layout.c128_row_nbytes == self.mla_bytes_per_token + from lightllm.common.state_cache_manager.deepseek4 import DeepseekV4StateCacheManager + + args = get_env_start_args() + big_page_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + self.big_page_buffers = DeepseekV4StateCacheManager( + size=_ceil_div(size, big_page_tokens), + layout=layout, + ) + logger.info( f"DeepseekV4MemoryManager pools: full_tokens={size} swa={self.swa_size}({self.swa_num_pages}p) " f"c4={self.c4_size}(L={self.n_c4}) c128={self.c128_size}(L={self.n_c128}) " @@ -680,8 +692,10 @@ def get_loadable_cpu_cache_end( loadable_end = min(requested_end, (loaded_start + token_capacity) // page * page) return loadable_end if loadable_end > loaded_start else 0 - def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4CpuCacheLoadPlan: - """Allocate a missing suffix and publish all derived mappings as one plan. + def prepare_cpu_cache_load( + self, token_num: int, loaded_end: int, resume_swa_slots: torch.Tensor + ) -> DeepseekV4CpuCacheLoadPlan: + """Allocate a missing history suffix using the request's continuation slots. ``loaded_end`` is the CPU checkpoint boundary. ``token_num`` may be smaller than the checkpoint page when a GPU radix prefix overlaps its @@ -702,32 +716,14 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C ) block_num = token_num // DSV4_PROMPT_CACHE_PAGE_SIZE - device = self.full_to_swa_indexs.device + device = self.swa_pool.buffer.device full_indexes_cpu = self.alloc(token_num) - swa_pages_cpu = self._alloc_swa_pages(2) mem_indexes = full_indexes_cpu.to(device, non_blocking=True) history_full_slots = mem_indexes.view(block_num, DSV4_PROMPT_CACHE_PAGE_SIZE) - resume_full_slots = history_full_slots[-1] - - swa_pages = swa_pages_cpu.to(device, non_blocking=True) - resume_swa_slots = ( - swa_pages[:, None] * DSV4_SWA_PAGE_SIZE - + torch.arange(DSV4_SWA_PAGE_SIZE, dtype=torch.int32, device=device)[None, :] - ).reshape(-1) - history_c4_slots = history_full_slots[:, 3::4] // 4 if self.n_c4 else None history_c128_slots = history_full_slots[:, 127::128] // 128 if self.n_c128 else None - resume_full_slots_long = resume_full_slots.long() - resume_swa_pages = swa_pages - resume_swa_page_deltas = torch.full( - resume_swa_pages.shape, - DSV4_SWA_PAGE_SIZE, - dtype=torch.int32, - device=device, - ) - return DeepseekV4CpuCacheLoadPlan( loaded_start=loaded_end - token_num, loaded_end=loaded_end, @@ -735,140 +731,15 @@ def prepare_cpu_cache_load(self, token_num: int, loaded_end: int) -> DeepseekV4C history_full_slots=history_full_slots, history_c4_slots=history_c4_slots, history_c128_slots=history_c128_slots, - resume_full_slots=resume_full_slots, resume_swa_slots=resume_swa_slots, - resume_full_slots_long=resume_full_slots_long, - resume_swa_pages=resume_swa_pages, - resume_swa_page_deltas=resume_swa_page_deltas, ) - def commit_cpu_cache_load_plan(self, plan: DeepseekV4CpuCacheLoadPlan) -> None: - """Publish an unpacked plan without allocating at the transaction boundary.""" - self.full_to_swa_indexs[plan.resume_full_slots_long] = plan.resume_swa_slots - self.swa_page_live_count.index_add_(0, plan.resume_swa_pages, plan.resume_swa_page_deltas) - return - - # ------------------------------------------------------------------ swa slot lifecycle - def register_swa_free_hook(self, fn) -> None: - """fn(need_pages): 在页 allocator 不足时尝试腾页(radix 对 ref==0 节点 free swa)。""" - self._free_radix_unreferenced_swa_fn = fn - return - def __getstate__(self): state = self.__dict__.copy() - # The radix tree is process-local; IPC readers only need its CUDA cache tensors. - state["_free_radix_unreferenced_swa_fn"] = None + # Pinned CPU checkpoints are process-local; IPC readers need only GPU storage. + state["big_page_buffers"] = None return state - def _alloc_swa_pages(self, need_pages: int) -> torch.Tensor: - if need_pages > self.swa_page_allocator.can_use_mem_size and self._free_radix_unreferenced_swa_fn is not None: - self._free_radix_unreferenced_swa_fn(need_pages - self.swa_page_allocator.can_use_mem_size) - return self.swa_page_allocator.alloc(need_pages) - - def _update_swa_page_counts(self, swa_slots: torch.Tensor, delta: int) -> torch.Tensor: - """按 slot 所在页更新存活计数,返回逐 slot 的页号。""" - pages = torch.div(swa_slots, DSV4_SWA_PAGE_SIZE, rounding_mode="floor") - ones = torch.full(pages.shape, delta, dtype=torch.int32, device=pages.device) - self.swa_page_live_count.index_add_(0, pages, ones) - return pages - - def alloc_swa_prefill( - self, - mem_indexes: torch.Tensor, - req_to_token_indexs: torch.Tensor, - req_list: List[int], - ready_list: List[int], - seq_list: List[int], - ) -> None: - """prefill prep: 为各请求位置 [ready, seq) 的新 token 分配位置对齐的 swa 槽。 - - 槽位不变式: slot(p) = page_base(p 所在页) + p%128,page_base % 128 == 0。 - 续页(start 非整页,只可能是首页)的 base 从上一 token 的映射派生 - (full_to_swa[req_to_token[req, start-1]],该 token 必在保留窗内);其余页全新分配。 - radix 命中(ready 必 128 对齐)的借用方从全新页开始,与节点持有页天然不相交。 - 当前 chunk 的 full 槽直接来自 generic preprocess 分配的 mem_indexes,因此不依赖 - req_to_token_indexs 已完成当前 chunk 的 scatter;只有续页的上一 token 查询旧 req 行。 - """ - page = DSV4_SWA_PAGE_SIZE - hold_req_id = self.max_request_num # padding 行的请求 id(req_manager.HOLD_REQUEST_ID) - - segs = [] # (req_idx, start, end, mem_offset, n_new_pages, has_cont_page) - total_new_pages = 0 - mem_offset = 0 - for req_idx, start, end in zip(req_list, ready_list, seq_list): - q_len = end - start - if req_idx == hold_req_id or end <= start: - mem_offset += q_len - continue - first_new_page = _ceil_div(start, page) - n_new = max(0, (end - 1) // page - first_new_page + 1) - segs.append((req_idx, start, end, mem_offset, n_new, start % page != 0)) - total_new_pages += n_new - mem_offset += q_len - if not segs: - return - - device = self.full_to_swa_indexs.device - mem_indexes = mem_indexes.reshape(-1) - new_pages = self._alloc_swa_pages(total_new_pages).to(device, non_blocking=True) if total_new_pages else None - page_cursor = 0 - for req_idx, start, end, mem_start, n_new, has_cont in segs: - positions = torch.arange(start, end, dtype=torch.int32, device=device) - page_local = torch.div(positions, page, rounding_mode="floor") - start // page - bases = torch.empty(((end - 1) // page - start // page + 1,), dtype=torch.int32, device=device) - if has_cont: - prev_slot = self.full_to_swa_indexs[req_to_token_indexs[req_idx, start - 1]] - bases[0] = prev_slot - (start - 1) % page - if n_new: - bases[1 if has_cont else 0 :] = new_pages[page_cursor : page_cursor + n_new] * page - page_cursor += n_new - slots = bases[page_local] + positions % page - full_slots = mem_indexes[mem_start : mem_start + end - start] - self.full_to_swa_indexs[full_slots] = slots - self._update_swa_page_counts(slots, 1) - return - - def alloc_swa_decode( - self, - req_list: List[int], - seq_list: List[int], - mem_indexes: torch.Tensor, - prev_full_indexes: torch.Tensor, - ) -> None: - """decode prep: 本步 token(位置 seq-1)的 swa 槽。整页起点开新页,否则上一 token 槽 +1 - (位置对齐不变式保证同页连续)。scatter 目标用当前步 mem_indexes。 - - 调用方传入每行前一 token 的 full 槽;MTP step>0 可直接使用同批前一列。""" - page = DSV4_SWA_PAGE_SIZE - hold_req_id = self.max_request_num - cont_rows, new_rows = [], [] - for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)): - if req_idx == hold_req_id or seq_len <= 0: - continue - if (seq_len - 1) % page == 0: - new_rows.append(i) - else: - cont_rows.append(i) - mem_indexes = mem_indexes.reshape(-1) - if cont_rows: - # Steady decode normally puts every request on the same page offset. - # Avoid Python-list indexing in that case: PyTorch copies the list to - # CUDA synchronously and waits for the previous decode graph. - all_rows = len(cont_rows) == len(req_list) - prev_full = prev_full_indexes.reshape(-1) if all_rows else prev_full_indexes.reshape(-1)[cont_rows] - prev_slots = self.full_to_swa_indexs[prev_full] - slots = prev_slots + 1 - dst_indexes = mem_indexes if all_rows else mem_indexes[cont_rows] - self.full_to_swa_indexs[dst_indexes] = slots - self._update_swa_page_counts(slots, 1) - if new_rows: - pages = self._alloc_swa_pages(len(new_rows)).to(self.full_to_swa_indexs.device, non_blocking=True) - slots = pages * page - dst_indexes = mem_indexes if len(new_rows) == len(req_list) else mem_indexes[new_rows] - self.full_to_swa_indexs[dst_indexes] = slots - self._update_swa_page_counts(slots, 1) - return - def alloc_dspark_swa_block(self, token_num: int, block_size: int): """Assign one temporary SWA scratch page to each DSpark proposal block. @@ -888,63 +759,23 @@ def alloc_dspark_swa_block(self, token_num: int, block_size: int): if req_num == 0: return ( torch.empty((0,), dtype=torch.int32, device="cpu"), - torch.empty((0,), dtype=torch.int32, device=self.full_to_swa_indexs.device), + torch.empty((0,), dtype=torch.int32, device=self.swa_pool.buffer.device), ) - device = self.full_to_swa_indexs.device - pages_cpu = self._alloc_swa_pages(req_num) + device = self.swa_pool.buffer.device + pages_cpu = self.swa_page_allocator.alloc(req_num) pages = pages_cpu.to(device, non_blocking=True) - # DSpark block 的物理槽由 attention index builder 直接从 page id - # 计算,不发布到 target 的全局 full->SWA 映射,也不参与 live count。 + # The proposal block addresses its own page directly, outside the request page table. return pages_cpu, pages def free_dspark_swa_block(self, pages_cpu: torch.Tensor) -> None: """Release only proposal scratch; token pages remain owned by the request.""" self.swa_page_allocator.free(pages_cpu) - def evict_swa(self, full_slots: torch.Tensor) -> None: - """回收 full 槽位对应的 swa 槽(出窗惰性回收 / free 级联 / 压力阀共用)。 - 未映射(-1)的槽位跳过;页计数减到 0 时整页归还 allocator。""" - if full_slots.numel() == 0: - return - full_slots = full_slots.to(self.full_to_swa_indexs.device, non_blocking=True).reshape(-1) - full_slots = torch.unique(full_slots[full_slots < self.size]) - if full_slots.numel() == 0: - return - swa_slots = self.full_to_swa_indexs[full_slots] - valid = swa_slots >= 0 - valid_slots = swa_slots[valid] - if valid_slots.numel() == 0: - return - self.full_to_swa_indexs[full_slots[valid]] = -1 - touched = torch.unique(self._update_swa_page_counts(valid_slots, -1)) - empty = touched[self.swa_page_live_count[touched] == 0] - if empty.numel() > 0: - self.swa_page_allocator.free(empty.to(torch.int32)) - return - - # ------------------------------------------------------------------ alloc/free (cascade) - def free(self, free_index: Union[torch.Tensor, List[int]]) -> None: - """释放 full token 槽位,级联回收其 swa 槽与 c4/c128 压缩槽。radix 驱逐、请求释放/暂停都走这里。 - - 先对 full 槽去重: 同批重复槽位会让映射 gather 出重复的压缩/swa 槽,导致 allocator 双重释放。""" - if isinstance(free_index, list): - free_index = torch.tensor(free_index, dtype=torch.int64) - if free_index.numel() > 0: - free_index = torch.unique(free_index) - self.evict_swa(free_index) - super().free(free_index) - return - def free_all(self): super().free_all() self.swa_page_allocator.free_all() - self.swa_page_live_count.zero_() - self.full_to_swa_indexs.fill_(-1) - self.full_to_swa_indexs[self.size :] = ( - self.swa_size + torch.arange(self.page_size, device=self.full_to_swa_indexs.device) % DSV4_SWA_PAGE_SIZE - ) - return + self.big_page_buffers.clear_to_init_state() # ------------------------------------------------------------------ packed codecs (torch reference) # 与 sglang/vllm 的 fp8_ds_mla 字节布局逐位对齐(ue8m0 幂次 scale)。这些 torch 实现是该 ABI 的 @@ -1012,48 +843,21 @@ def _unpack_indexer_k(self, packed: torch.Tensor) -> torch.Tensor: return (k_fp8 * scale).to(self.dtype) # ------------------------------------------------------------------ cache write paths - def pack_mla_kv_to_cache(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - """标准 operator 写入路径。要求本步已对 mem_index 调过 ``alloc_swa``(prep 阶段); - HOLD/padding 槽位映射到 swa HOLD 槽,写入无害。""" - if kv.shape[0] == 0: - return - from lightllm.models.deepseek_v4.triton_kernel.destindex_copy_kv_flashmla_dsv4 import ( - destindex_copy_kv_flashmla_dsv4, - ) - - swa_slots = self.full_to_swa_indexs[mem_index.cuda().long().reshape(-1)] - destindex_copy_kv_flashmla_dsv4( - kv.reshape(-1, self.mla_head_dim), - swa_slots, - self.swa_pool.get_layer_buffer(layer_index), - self.swa_pool.page_size, - ) - return - def pack_mla_kv_to_cache_fused_norm_rope( self, layer_index: int, - mem_index: torch.Tensor, + swa_slots: torch.Tensor, kv: torch.Tensor, kv_weight: torch.Tensor, eps: float, freqs_cis: torch.Tensor, positions: torch.Tensor, - swa_slots: Optional[torch.Tensor] = None, ): - """同 pack_mla_kv_to_cache,但 rmsnorm + 尾部交错 rope 融合进写入 kernel - 并省掉 bf16 kv 中间量。kv 为 wkv 投影原始输出 [T, head_dim+rope_dim]。""" + """Fuse RMSNorm and RoPE into the packed writer at request-owned SWA slots.""" from lightllm.models.deepseek_v4.triton_kernel.norm_rope_cuda import ( fused_k_norm_rope_flashmla, ) - if swa_slots is None: - swa_slots = self.full_to_swa_indexs[mem_index.reshape(-1)] - swa_slots = torch.where( - swa_slots < 0, - torch.full_like(swa_slots, self.swa_pool.HOLD_TOKEN_MEMINDEX), - swa_slots, - ) fused_k_norm_rope_flashmla( kv=kv, kv_weight=kv_weight, diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 0cd5401de8..92034d671f 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -98,19 +98,15 @@ def copy_mem_to_mem(self, src_mem_index: torch.Tensor, dst_mem_index: torch.Tens pool.buffer.index_copy_(1, dst_pages, pool.buffer.index_select(1, src_pages)) def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( - DeepseekV4MemoryManager, - ) - - mem_manager: DeepseekV4MemoryManager = self.mem_manager - mem_manager.pack_mla_kv_to_cache(layer_index, mem_index, kv) - return + raise NotImplementedError("DeepSeek-V4 writes packed KV using request-owned SWA slots") - def pack_cpu_cache_pages(self, source_mem_indexes: torch.Tensor, staging: torch.Tensor) -> None: + def pack_cpu_cache_pages( + self, source_mem_indexes: torch.Tensor, source_req_meta: torch.Tensor, staging: torch.Tensor + ) -> None: """Pack complete DS4 checkpoints into caller-owned CUDA staging.""" from lightllm.models.deepseek_v4.triton_kernel.cpu_cache_io import pack_gpu_cache_to_staging - pack_gpu_cache_to_staging(self.mem_manager, source_mem_indexes, staging) + pack_gpu_cache_to_staging(self.mem_manager, source_mem_indexes, source_req_meta, staging) return def scatter_packed_cpu_cache_pages( 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 71158ac97a..b3d1b7cfcd 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py +++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py @@ -95,6 +95,7 @@ def load_cpu_cache_to_gpu( g_infer_context.req_manager.restore_big_page_state( big_page_buffer_idx=big_page_buffer_ids_cpu[-1], req=req, + checkpoint_len=req.cur_kv_len, ) return diff --git a/lightllm/common/req_manager/__init__.py b/lightllm/common/req_manager/__init__.py index cf5ebdc039..611f45b599 100644 --- a/lightllm/common/req_manager/__init__.py +++ b/lightllm/common/req_manager/__init__.py @@ -5,4 +5,4 @@ __all__ = ["ReqManager", "HybridAttentionReqManager", "ReqManagerForMamba", "ReqSamplingParamsManager"] -from .deepseek4 import DeepseekV4ReqManager, DeepseekV4PromptCachePayload, DeepseekV4PromptCacheValueOps +from .deepseek4 import DeepseekV4ReqManager diff --git a/lightllm/common/req_manager/deepseek4.py b/lightllm/common/req_manager/deepseek4.py index 4176db4267..58bcfaa6b5 100644 --- a/lightllm/common/req_manager/deepseek4.py +++ b/lightllm/common/req_manager/deepseek4.py @@ -1,81 +1,22 @@ -from dataclasses import dataclass -from typing import List, Optional, TYPE_CHECKING +from typing import Optional import torch -from .base import ReqManager +from .hybrid_base import HybridAttentionReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_C4_PAGE_SIZE, DSV4_PROMPT_CACHE_PAGE_SIZE -from lightllm.utils.envs_utils import get_env_start_args -from lightllm.utils.log_utils import init_logger -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_SWA_PAGE_SIZE, + DSV4_PROMPT_CACHE_PAGE_SIZE, +) +from lightllm.common.state_cache_manager.deepseek4 import DeepseekV4StateCacheManager -if TYPE_CHECKING: - from lightllm.server.router.model_infer.infer_batch import InferReq -logger = init_logger(__name__) +class DeepseekV4ReqManager(HybridAttentionReqManager): + """Own request-private SWA pages and restore aligned continuation checkpoints. - -@dataclass -class DeepseekV4PromptCachePayload: - """prompt cache 载荷: swa 按页有效性 bitmap 和最后有效页。 - - 压缩历史由统一 token 页持有;SWA 通过 full slot 映射定位。 - 槽位与 compressor 状态都不进载荷;c4 compressor 状态随 swa 页 - 生灭。c128 状态按 request ring 寻址;prompt cache 的 256-token 边界同时是 c128 - 分组边界,命中后新分组会在首次读取前覆写完整 128-token 窗口,因而无需保存状态。 - - * ``swa_page_valid``: cpu bool [cache_len // page],插入时按当下 full_to_swa 映射写定 - (页内 token 映射全有效才为 True)。匹配层据此把命中裁剪到"结尾页有效"的 page 边界, - swa 压力阀回收节点页时清零。""" - - cache_len: int - swa_page_valid: Optional[torch.Tensor] = None - swa_last_valid_page: int = -1 - - def refresh_swa_last_valid_page(self) -> None: - if self.swa_page_valid is None: - self.swa_last_valid_page = -1 - return - valid_idx = torch.nonzero(self.swa_page_valid).flatten() - self.swa_last_valid_page = -1 if valid_idx.numel() == 0 else int(valid_idx[-1].item()) - return - - def valid_match_length(self, natural_len: int, page: int) -> int: - if self.swa_last_valid_page < 0: - return 0 - return (int(self.swa_last_valid_page) + 1) * page - - -class DeepseekV4PromptCacheValueOps: - def __init__(self, req_manager: "DeepseekV4ReqManager"): - self.req_manager = req_manager - - def slice(self, payload: DeepseekV4PromptCachePayload, start: int, end: int): - return self.req_manager.slice_prompt_cache_payload(payload, start, end) - - def concat(self, payloads: List[DeepseekV4PromptCachePayload]): - return self.req_manager.concat_prompt_cache_payloads(payloads) - - def free(self, payload: DeepseekV4PromptCachePayload): - # 槽位资源全部由 mem_manager.free(full_slots) 级联回收,载荷本身没有需要释放的资源。 - return - - def valid_match_length(self, payload: Optional[DeepseekV4PromptCachePayload], natural_len: int) -> int: - """radix 匹配裁剪: 返回 <= natural_len 的最大 prompt-cache 边界 L',使结尾页有效。 - - 有效性可能非单调(owner 生前从左驱逐、后续阀从尾回收),中段 invalid 页不挡更 - 靠后的有效命中(注意力只回看最后一个窗口)。""" - if payload is None: - return 0 - return payload.valid_match_length(natural_len, self.req_manager.get_prompt_cache_page_size()) - - -class DeepseekV4ReqManager(ReqManager): - """DeepSeek-V4 的请求级管理。 - - 负责 req/seq/MTP 布局、SWA 回收水位线和派生槽位准备;具体池结构、映射和分配器 - 由 ``DeepseekV4MemoryManager`` 持有。对象先于 mem manager 创建,模型初始化后再接入。 + The page table covers the whole sequence, so a prefill chunk can keep all of + its new KV until attention completes. Only pages preceding the next chunk's + retained window are recycled. Speculative retries reuse their held pages. """ def __init__( @@ -83,343 +24,136 @@ def __init__( max_request_num, max_sequence_length, mem_manager: Optional[DeepseekV4MemoryManager] = None, - sliding_window: Optional[int] = None, + sliding_window=None, ): - super().__init__(max_request_num, max_sequence_length, mem_manager) - + super().__init__(max_request_num, max_sequence_length, None) self.sliding_window = sliding_window - # 出窗回收水位线: -1 表示该 req 尚未见过 prefill chunk(首个 chunk 的 ready_cache_len - # 即共享前缀边界,作为永不下探的回收下界)。 - self._swa_evict_marks = [-1 for _ in range(max_request_num + 1)] - self._swa_allocated_end = [0 for _ in range(max_request_num + 1)] - return - - # ------------------------------------------------------------------ swa slot prep (per step) - def _swa_retain_len(self) -> int: - """出窗回收的保留长度 = window + 一个 radix 页。 - - 多留一页使「最近一个完成的 prompt-cache 边界」的结尾页恒驻留: 若回收只留 window, - 则任何非对齐时刻该边界的结尾页都已被部分回收,插入门会把所有插入裁到 0。 - V4 prompt-cache 页取 256 token,正好覆盖一个 c4 物理页对应的 token 范围。""" - return int(self.sliding_window) + self.get_prompt_cache_page_size() - - def _align_swa_evict_frontier(self, raw_frontier: int) -> int: - """SWA 回收水位线按 prompt-cache 页向下对齐。 - - bitmap 的有效性是 prompt-cache page 粒度;若水位线切进页面中间,该页会被判为 - invalid,即使靠近命中边界的窗口实际仍完整驻留。""" - page = self.get_prompt_cache_page_size() - raw_frontier = max(0, int(raw_frontier)) - return raw_frontier // page * page - - def prepare_prefill_swa( - self, - req_list: List[int], - ready_list: List[int], - seq_list: List[int], - mem_indexes: torch.Tensor, - ) -> None: - """prefill prep: 为本 chunk 全部新 token(位置 [ready, seq))分配位置对齐的 swa 槽, - 并回收已出窗位置的槽。 - - 本 chunk 起点 L = ready_cache_len,首个新 token(位置 L)的窗口是 [L-W+1, L];回收 - 边界再额外保留一个 radix 页(_swa_retain_len),即位置 < L-retain+1。先回收再分配。 - 当前 chunk 的 full slots 直接使用 generic preprocess 分配的 mem_indexes,因而可以 - 在通用 req_to_token scatter 之前执行。""" - self.mem_manager: DeepseekV4MemoryManager - if self.sliding_window is not None: - retain = self._swa_retain_len() - evict_slots = [] - for req_idx, ready_len in zip(req_list, ready_list): - if req_idx == self.HOLD_REQUEST_ID: - continue - mark = self._swa_evict_marks[req_idx] - if mark < 0: - # 首个 chunk: [0, ready_len) 是 radix 共享前缀,其 swa 槽归 radix 所有,不可回收。 - self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(ready_len) - continue - evict_end = self._align_swa_evict_frontier(ready_len - retain + 1) - if evict_end > mark: - evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) - self._swa_evict_marks[req_idx] = evict_end - if evict_slots: - self.mem_manager.evict_swa(torch.cat(evict_slots)) - self.mem_manager.alloc_swa_prefill( - mem_indexes, - self.req_to_token_indexs, - req_list=req_list, - ready_list=ready_list, - seq_list=seq_list, - ) - for req_idx, seq_len in zip(req_list, seq_list): - self._swa_allocated_end[req_idx] = seq_len - return - - def prepare_decode( - self, - b_req_idx_cpu, - b_seq_len_cpu, - b_mtp_index_cpu, - mem_indexes, - ): - """decode 每步槽位 prep。在 BaseModel 的通用 req scatter 与 attention metadata - 构建前调用;DeepSeek-V4 MTP draft layer 只需要 SWA 槽位。""" - max_mtp_index = int(b_mtp_index_cpu.max().item()) - steps = range(max_mtp_index + 1) - - batch_size = b_mtp_index_cpu.shape[0] - slots_per_req = max_mtp_index + 1 - assert batch_size % slots_per_req == 0 - req_list = b_req_idx_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() - mem_indexes_by_req = mem_indexes.reshape(-1, slots_per_req) - for step in steps: - step_req_list = req_list[step::slots_per_req] - step_seq_list = seq_list[step::slots_per_req] - self.prepare_decode_swa( - step_req_list, - step_seq_list, - mem_indexes_by_req[:, step], - prev_mem_indexes=mem_indexes_by_req[:, step - 1] if step > 0 else None, - ) - return - - def prepare_prefill( - self, - b_req_idx_cpu: torch.Tensor, - b_ready_cache_len_cpu: torch.Tensor, - b_seq_len_cpu: torch.Tensor, - mem_indexes: torch.Tensor, - ) -> None: - """prefill 槽位 prep: 直接消费 generic preprocess 分配的 full slots,在 - BaseModel 的通用 req scatter 与 attention metadata 构建之前完成。""" - req_list = b_req_idx_cpu.tolist() - ready_list = b_ready_cache_len_cpu.tolist() - seq_list = b_seq_len_cpu.tolist() - mem_indexes = mem_indexes.reshape(-1) - self.prepare_prefill_swa( - req_list=req_list, - ready_list=ready_list, - seq_list=seq_list, - mem_indexes=mem_indexes, + self.req_to_swa_pages = torch.full( + (max_request_num + 1, self.req_to_token_indexs.shape[1] // DSV4_SWA_PAGE_SIZE), + -1, + dtype=torch.int32, + device="cuda", ) - return - - def prepare_pd_decode_cache( - self, - req_list: List[int], - ready_list: List[int], - seq_list: List[int], - new_full_slots: torch.Tensor, - ) -> None: - """Allocate DSV4 derived slots for request-major suffixes received from peers.""" - page = self.get_prompt_cache_page_size() - assert len(req_list) == len(ready_list) == len(seq_list) and len(req_list) > 0 - assert all(ready % page == 0 and seq_len > ready for ready, seq_len in zip(ready_list, seq_list)) - assert new_full_slots.numel() == sum(seq_len - ready for ready, seq_len in zip(ready_list, seq_list)) - - new_full_slots = new_full_slots.reshape(-1).to(self.req_to_token_indexs.device, non_blocking=True) - - # swa 只保存最后一部分,前面的不需要 - swa_start_list = [] - swa_parts = [] - offset = 0 - # swa 不一样 - for ready, seq_len in zip(ready_list, seq_list): - swa_start = max(ready, max(0, seq_len // page * page - page)) - swa_start_list.append(swa_start) - suffix_len = seq_len - ready - swa_parts.append(new_full_slots[offset + swa_start - ready : offset + suffix_len]) - offset += suffix_len - swa_full_slots = swa_parts[0] if len(swa_parts) == 1 else torch.cat(swa_parts) - - self.mem_manager.alloc_swa_prefill( - swa_full_slots, - self.req_to_token_indexs, - req_list=req_list, - ready_list=swa_start_list, - seq_list=seq_list, + self._swa_pages = [{} for _ in range(max_request_num)] + if mem_manager is not None: + self.bind_mem_manager(mem_manager) + + def bind_mem_manager(self, mem_manager): + super().bind_mem_manager(mem_manager) + self.req_to_swa_pages[self.HOLD_REQUEST_ID].fill_(mem_manager.swa_num_pages) + # IPC pack/unpack workers need the GPU table, without the CPU owner lists. + mem_manager.req_to_swa_pages = self.req_to_swa_pages + + def get_swa_page_need(self, req_idx, start, end): + pages = self._swa_pages[req_idx] + first_retained = max(0, start - self.sliding_window - DSV4_PROMPT_CACHE_PAGE_SIZE + 1) // DSV4_SWA_PAGE_SIZE + missing = sum( + p not in pages + for p in range(start // DSV4_SWA_PAGE_SIZE, (end + DSV4_SWA_PAGE_SIZE - 1) // DSV4_SWA_PAGE_SIZE) ) - for req_idx, swa_start, seq_len in zip(req_list, swa_start_list, seq_list): - self._swa_evict_marks[req_idx] = swa_start - self._swa_allocated_end[req_idx] = seq_len - return + released = sum(p < first_retained for p in pages) + return max(0, missing - released) - def prepare_decode_swa( - self, - req_list: List[int], - seq_list: List[int], - mem_indexes: torch.Tensor, - prev_mem_indexes: Optional[torch.Tensor] = None, - ) -> None: - """decode prep: 回收出窗槽并为本步新 token 分配位置对齐的 swa 槽。当前 query 位置 - seq_len-1 的窗口是 [seq_len-W, seq_len-1];回收边界额外保留一个 radix 页 - (_swa_retain_len),即位置 < seq_len-retain。先回收再分配。 - seq_len/req_idx 从 CPU 镜像读(host 算术,无 D2H);水位线 _swa_evict_marks 仍是 host 状态。""" - assert self.mem_manager is not None - if self.sliding_window is not None: - retain = self._swa_retain_len() - evict_slots = [] - for req_idx, seq_len in zip(req_list, seq_list): - if req_idx == self.HOLD_REQUEST_ID: - continue - mark = self._swa_evict_marks[req_idx] - if mark < 0: - # direct-decode 中 [0, seq_len-1) 是已有 KV;exact hit 时这段前缀归 radix 所有, - # 水位必须从前缀末端开始,不能由请求回收。 - self._swa_evict_marks[req_idx] = self._align_swa_evict_frontier(seq_len - 1) - continue - evict_end = self._align_swa_evict_frontier(seq_len - retain) - if evict_end > mark: - evict_slots.append(self.req_to_token_indexs[req_idx, mark:evict_end]) - self._swa_evict_marks[req_idx] = evict_end - if evict_slots: - self.mem_manager.evict_swa(torch.cat(evict_slots)) - new_rows = [ - i - for i, (req_idx, seq_len) in enumerate(zip(req_list, seq_list)) - if req_idx != self.HOLD_REQUEST_ID and seq_len > self._swa_allocated_end[req_idx] - ] - if not new_rows: + def prepare_swa(self, req_idx, start, end): + if req_idx == self.HOLD_REQUEST_ID: return - req_list = [req_list[i] for i in new_rows] - seq_list = [seq_list[i] for i in new_rows] - rows = torch.tensor(new_rows, dtype=torch.int64, device=mem_indexes.device) - mem_indexes = mem_indexes[rows] - if prev_mem_indexes is not None: - prev_mem_indexes = prev_mem_indexes[rows] - if prev_mem_indexes is None: - prev_meta = g_pin_mem_manager.gen_from_list( - key="dsv4_swa_decode_prev", - data=[x for req_idx, seq_len in zip(req_list, seq_list) for x in (req_idx, seq_len - 2)], - dtype=torch.int64, - ).to(self.req_to_token_indexs.device, non_blocking=True) - prev_meta = prev_meta.view(-1, 2) - prev_mem_indexes = self.req_to_token_indexs[prev_meta[:, 0], prev_meta[:, 1]] - self.mem_manager.alloc_swa_decode( - req_list, - seq_list, - mem_indexes, - prev_mem_indexes, - ) - for req_idx, seq_len in zip(req_list, seq_list): - self._swa_allocated_end[req_idx] = seq_len - return - - def init_compress_state(self, req_idx: int): - """新请求开始时重置 runtime 水位线(对应 mamba 的 init_linear_att_state 调用点)。 - - c4 状态随 swa 页寻址;c128 request ring 依靠 overwrite-before-read,不做大块清零。""" - self.clear_runtime_state(req_idx) - return - - def finish_cpu_cache_load(self, req_idx: int, loaded_len: int) -> None: - """Keep only the final restored 256-token SWA page eligible for radix reuse.""" - self._swa_evict_marks[req_idx] = loaded_len - self.get_prompt_cache_page_size() - self._swa_allocated_end[req_idx] = loaded_len - return - - def alloc(self): - req_idx = super().alloc() - if req_idx is not None: - self.init_compress_state(req_idx) - return req_idx - - def clear_runtime_state(self, req_idx: int): - # swa 槽位本身由 mem_manager.free 级联回收(随 full 槽位),这里只复位出窗水位线。 - self._swa_evict_marks[req_idx] = -1 - self._swa_allocated_end[req_idx] = 0 - return - - def get_prompt_cache_value_ops(self): - return DeepseekV4PromptCacheValueOps(self) + page_size = DSV4_SWA_PAGE_SIZE + pages = self._swa_pages[req_idx] + retain = self.sliding_window + DSV4_PROMPT_CACHE_PAGE_SIZE + first_retained_page = max(0, start - retain + 1) // page_size + evicted = [position for position in pages if position < first_retained_page] + if evicted: + self.mem_manager.swa_page_allocator.free(torch.tensor([pages.pop(p) for p in evicted], dtype=torch.int32)) + self.req_to_swa_pages[req_idx, :first_retained_page] = -1 + missing = [p for p in range(start // page_size, (end + page_size - 1) // page_size) if p not in pages] + if missing: + allocated = self.mem_manager.swa_page_allocator.alloc(len(missing)) + pages.update(zip(missing, allocated.tolist())) + self.req_to_swa_pages[req_idx, missing] = allocated.to(device="cuda", non_blocking=True) def get_prompt_cache_page_size(self): return DSV4_PROMPT_CACHE_PAGE_SIZE - def compute_swa_page_valid(self, full_slots: torch.Tensor) -> torch.Tensor: - """按当下 full_to_swa 映射给出按页有效性: full_slots [L](L 为 page 整数倍) -> - cpu bool [L/page],页内全部映射有效才为 True。GPU gather + 同步,测试/校验用; - 插入热路径用 swa_page_valid_from_watermark(纯 CPU,免同步)。""" - page = self.get_prompt_cache_page_size() - assert full_slots.numel() % page == 0 - if full_slots.numel() == 0: - return torch.zeros((0,), dtype=torch.bool) - swa = self.mem_manager.full_to_swa_indexs[full_slots.cuda().long().reshape(-1)] - return (swa.view(-1, page) >= 0).all(dim=1).cpu() - - def swa_page_valid_from_watermark(self, req_idx: int, cache_len: int) -> torch.Tensor: - """插入时的按页有效性,纯 CPU: 请求自有 token 的 swa 映射只被出窗水位线回收 - (阀不触活跃请求,级联只在 free 时),页 p 全驻留 ⟺ 页起点 page*p >= 水位线。 - - 与 compute_swa_page_valid 在插入时刻对自有 token 等价,但不做 GPU gather/同步—— - router 关键路径上每次插入省一次对全部在途 kernel 的等待。bitmap 中借入前缀 - ([0, ready) 的页)的行在 radix insert 切片时被丢弃(既有节点保留自己的 bitmap), - 其取值无影响。""" - page = self.get_prompt_cache_page_size() - mark = max(0, self._swa_evict_marks[req_idx]) - n_pages = int(cache_len) // page - return torch.arange(n_pages, dtype=torch.long) * page >= mark - - def slice_prompt_cache_payload(self, payload: DeepseekV4PromptCachePayload, start: int, end: int): - start = int(start) - end = int(end) - page = self.get_prompt_cache_page_size() - # radix page 保证分裂点页对齐,bitmap 可整页切分。 - ans = DeepseekV4PromptCachePayload( - cache_len=end - start, - swa_page_valid=payload.swa_page_valid[start // page : end // page].clone() - if payload.swa_page_valid is not None - else None, - ) - ans.refresh_swa_last_valid_page() - return ans - - def concat_prompt_cache_payloads(self, payloads: List[DeepseekV4PromptCachePayload]): - if len(payloads) == 0: - return None - bitmaps = [p.swa_page_valid for p in payloads] - ans = DeepseekV4PromptCachePayload( - cache_len=sum(p.cache_len for p in payloads), - swa_page_valid=torch.cat(bitmaps, dim=0) if all(b is not None for b in bitmaps) else None, + def get_swa_slots(self, req_idx, positions): + return ( + self.req_to_swa_pages[req_idx, positions // DSV4_SWA_PAGE_SIZE] * DSV4_SWA_PAGE_SIZE + + positions % DSV4_SWA_PAGE_SIZE ) - if ans.swa_page_valid is None: - return ans - page = self.get_prompt_cache_page_size() - page_offset = 0 - last_valid_page = -1 - for item in payloads: - item_last = int(getattr(item, "swa_last_valid_page", -1)) - if item_last >= 0: - last_valid_page = page_offset + item_last - page_offset += int(item.cache_len) // page - ans.swa_last_valid_page = last_valid_page - return ans - - def build_prompt_cache_payload( - self, - cache_len: int, - ) -> DeepseekV4PromptCachePayload: - """构造插入载荷。compressor 状态不进载荷(c4 随 swa 页生灭、c128 在 256 对齐 - 恢复点依靠 overwrite-before-read),cache_len 不再受序列末端约束。 - swa_page_valid 不在此填: 它必须用插入时刻的映射(infer batch 在 insert 前补)。""" - assert self.mem_manager is not None - return DeepseekV4PromptCachePayload(cache_len=int(cache_len)) + def prepare_prefill(self, b_req_idx_cpu, b_ready_cache_len_cpu, b_seq_len_cpu): + for req_idx, start, end in zip(b_req_idx_cpu.tolist(), b_ready_cache_len_cpu.tolist(), b_seq_len_cpu.tolist()): + self.prepare_swa(req_idx, start, end) + + def prepare_decode(self, b_req_idx_cpu, b_seq_len_cpu, b_mtp_index_cpu): + # Verification rows are request-major and include consecutive MTP positions. + width = int(b_mtp_index_cpu.max().item()) + 1 + reqs, seqs = b_req_idx_cpu.tolist(), b_seq_len_cpu.tolist() + for i in range(0, len(reqs), width): + self.prepare_swa(reqs[i], seqs[i] - 1, seqs[i + width - 1]) + + def prepare_pd_decode_cache(self, req_list, seq_list): + for req_idx, end in zip(req_list, seq_list): + # PD may end inside a compression group; preserve the complete tail + # required by its SWA and C4 continuation layout. + start = max( + 0, (end - 1) // DSV4_PROMPT_CACHE_PAGE_SIZE * DSV4_PROMPT_CACHE_PAGE_SIZE - DSV4_PROMPT_CACHE_PAGE_SIZE + ) + self.prepare_swa(req_idx, start, end) + + def create_small_page_cache_manager(self, size): + self.small_page_buffers = DeepseekV4StateCacheManager(size, self.mem_manager.cpu_cache_layout) + return self.small_page_buffers + + def init_hybrid_attention_state(self, req): + self.clear_runtime_state(req.req_idx) + + def save_state(self, req_idx, buffer_idx, state_cache_manager, checkpoint_len): + assert checkpoint_len > 0 and checkpoint_len % DSV4_PROMPT_CACHE_PAGE_SIZE == 0 + manager = self.mem_manager + end_page = checkpoint_len // DSV4_SWA_PAGE_SIZE + pages = self.req_to_swa_pages[req_idx, end_page - 2 : end_page].long() + swa, c4, indexer = state_cache_manager.get_state_cache(buffer_idx) + swa.copy_(manager.swa_pool.buffer.index_select(1, pages), non_blocking=True) + if manager.n_c4: + slots = pages[-1] * DSV4_SWA_PAGE_SIZE + torch.arange(124, 128, device="cuda") + rows = (slots // DSV4_SWA_PAGE_SIZE) * manager.c4_state_ring + slots % manager.c4_state_ring + c4.copy_(manager.c4_state_buffer.index_select(1, rows), non_blocking=True) + indexer.copy_(manager.c4_indexer_state_buffer.index_select(1, rows), non_blocking=True) + + def restore_state(self, req, state_cache_manager, buffer_idx, checkpoint_len): + assert checkpoint_len > 0 and checkpoint_len % DSV4_PROMPT_CACHE_PAGE_SIZE == 0 + self.clear_runtime_state(req.req_idx) + self.prepare_swa(req.req_idx, checkpoint_len - DSV4_PROMPT_CACHE_PAGE_SIZE, checkpoint_len) + manager = self.mem_manager + end_page = checkpoint_len // DSV4_SWA_PAGE_SIZE + pages = self.req_to_swa_pages[req.req_idx, end_page - 2 : end_page].long() + swa, c4, indexer = state_cache_manager.get_state_cache(buffer_idx) + manager.swa_pool.buffer.index_copy_(1, pages, swa.cuda(non_blocking=True)) + if manager.n_c4: + slots = pages[-1] * DSV4_SWA_PAGE_SIZE + torch.arange(124, 128, device="cuda") + rows = (slots // DSV4_SWA_PAGE_SIZE) * manager.c4_state_ring + slots % manager.c4_state_ring + manager.c4_state_buffer.index_copy_(1, rows, c4.cuda(non_blocking=True)) + manager.c4_indexer_state_buffer.index_copy_(1, rows, indexer.cuda(non_blocking=True)) + # A 256-token checkpoint closes both compressor groups. The next C128 + # group overwrites all of its rows before reading them. + + def clear_runtime_state(self, req_idx): + pages = self._swa_pages[req_idx] + if pages: + self.mem_manager.swa_page_allocator.free(torch.tensor(list(pages.values()), dtype=torch.int32)) + pages.clear() + self.req_to_swa_pages[req_idx].fill_(-1) def free(self, free_req_indexes, free_token_index): - """dense/swa/压缩槽全部经 mem_manager.free(free_token_index) 级联回收。""" - for req_index in free_req_indexes: - self.clear_runtime_state(req_index) + for req_idx in free_req_indexes: + self.clear_runtime_state(req_idx) super().free(free_req_indexes, free_token_index) - return - def free_req(self, free_req_index: int): + def free_req(self, free_req_index): self.clear_runtime_state(free_req_index) - return super().free_req(free_req_index) + super().free_req(free_req_index) def free_all(self): + for req_idx in range(self.max_request_num): + self.clear_runtime_state(req_idx) super().free_all() - self._swa_evict_marks = [-1 for _ in range(self.max_request_num + 1)] - self._swa_allocated_end = [0 for _ in range(self.max_request_num + 1)] - return diff --git a/lightllm/common/req_manager/hybrid_base.py b/lightllm/common/req_manager/hybrid_base.py index 948313116d..3715407516 100644 --- a/lightllm/common/req_manager/hybrid_base.py +++ b/lightllm/common/req_manager/hybrid_base.py @@ -39,30 +39,32 @@ def create_small_page_cache_manager(self, size: int): def init_hybrid_attention_state(self, req: "InferReq"): """无前缀缓存命中时,初始化已分配请求槽位的 GPU 运行态。""" - def restore_big_page_state(self, big_page_buffer_idx: int, req: "InferReq"): + def restore_big_page_state(self, big_page_buffer_idx: int, req: "InferReq", checkpoint_len: int): """将指定大页槽位的 CPU checkpoint 恢复到请求 GPU 运行态。""" - self.restore_state(req, self.big_page_buffers, big_page_buffer_idx) + self.restore_state(req, self.big_page_buffers, big_page_buffer_idx, checkpoint_len) - def restore_small_page_state(self, req: "InferReq"): + def restore_small_page_state(self, req: "InferReq", checkpoint_len: int): """将 req.shared_kv_node 对应的小页 checkpoint 恢复到请求 GPU 运行态。""" - self.restore_state(req, self.small_page_buffers, req.shared_kv_node.small_page_buffer_idx) + self.restore_state(req, self.small_page_buffers, req.shared_kv_node.small_page_buffer_idx, checkpoint_len) @abstractmethod - def restore_state(self, req: "InferReq", state_cache_manager, buffer_idx: int): + def restore_state(self, req: "InferReq", state_cache_manager, buffer_idx: int, checkpoint_len: int): """CPU checkpoint → 请求 GPU 运行态;大小页共用,不负责前缀匹配或 full KV 索引恢复。""" - def save_big_page_states(self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int]): + def save_big_page_states( + self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int], checkpoint_lens: List[int] + ): """批量保存请求 GPU 运行态到已分配的大页槽位,buffer_indexes 中的 -1 表示跳过。 b_req_idx 与 req_indexes 分别为同一批请求的 GPU 索引张量和 CPU 索引列表。 默认逐请求保存,模型可覆盖为批量拷贝算子。 """ - for req_idx, buffer_idx in zip(req_indexes, buffer_indexes): + for req_idx, buffer_idx, length in zip(req_indexes, buffer_indexes, checkpoint_lens): if buffer_idx != -1: - self.save_state(req_idx, buffer_idx, self.big_page_buffers) + self.save_state(req_idx, buffer_idx, self.big_page_buffers, length) @abstractmethod - def save_state(self, req_idx: int, buffer_idx: int, state_cache_manager): + def save_state(self, req_idx: int, buffer_idx: int, state_cache_manager, checkpoint_len: int): """请求 GPU 运行态 → 指定 CPU checkpoint 槽位;大小页共用,调用方负责分配槽位。""" def update_mtp_state(self, b_req_mtp_start_loc, b_req_idx, b_mtp_index, accepted_index, verify_width): diff --git a/lightllm/common/req_manager/linear_att.py b/lightllm/common/req_manager/linear_att.py index ccab8a06f6..25dcff9ce2 100644 --- a/lightllm/common/req_manager/linear_att.py +++ b/lightllm/common/req_manager/linear_att.py @@ -63,7 +63,9 @@ def create_small_page_cache_manager(self, size: int): self.small_page_buffers = LinearAttCacheManager(size=size, linear_config=self.linear_config) return self.small_page_buffers - def save_big_page_states(self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int]): + def save_big_page_states( + self, b_req_idx: torch.Tensor, req_indexes: List[int], buffer_indexes: List[int], checkpoint_lens: List[int] + ): from lightllm.common.basemodel.triton_kernel.linear_att_copy import copy_linear_att_state_to_kv_buffer buffer_indexes = torch.tensor(buffer_indexes, dtype=torch.int32, device="cpu").cuda(non_blocking=True) @@ -79,7 +81,9 @@ def save_big_page_states(self, b_req_idx: torch.Tensor, req_indexes: List[int], ) return - def save_state(self, req_idx: int, buffer_idx: int, state_cache_manager: LinearAttCacheManager): + def save_state( + self, req_idx: int, buffer_idx: int, state_cache_manager: LinearAttCacheManager, checkpoint_len: int + ): # checkpoint 只保存标准 conv 窗口和请求的基准 SSM 状态,不包含 MTP 扩展运行态。 conv_cache_width = self.linear_config.get_conv_state_shape()[-1] gpu_conv_state = self.req_to_conv_state.buffer[:, req_idx, ..., :conv_cache_width] @@ -109,7 +113,9 @@ def update_mtp_state(self, b_req_mtp_start_loc, b_req_idx, b_mtp_index, accepted verify_width=verify_width, ) - def restore_state(self, req: "InferReq", state_cache_manager: LinearAttCacheManager, buffer_idx: int): + def restore_state( + self, req: "InferReq", state_cache_manager: LinearAttCacheManager, buffer_idx: int, checkpoint_len: int + ): conv_state, ssm_state = state_cache_manager.get_state_cache(buffer_idx=buffer_idx) conv_dest = req.req_idx ssm_dest = req.req_idx * (self.mtp_step + 1) diff --git a/lightllm/common/state_cache_manager/__init__.py b/lightllm/common/state_cache_manager/__init__.py index af0479a158..88d73f5e2c 100644 --- a/lightllm/common/state_cache_manager/__init__.py +++ b/lightllm/common/state_cache_manager/__init__.py @@ -1,13 +1,18 @@ from .base import StateCacheManager from .layer_cache import LayerCache from .linear_att import LinearAttCacheConfig, LinearAttCacheManager +from .deepseek4 import DeepseekV4StateCacheManager def get_hybrid_cache_config(): """Return the model-specific layout used by hybrid CPU/disk cache pages.""" - from lightllm.utils.config_utils import is_linear_att_mixed_model + from lightllm.utils.config_utils import is_linear_att_mixed_model, get_model_type from lightllm.utils.envs_utils import get_env_start_args + if get_model_type(get_env_start_args().model_dir) == "deepseek_v4": + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4CpuCacheLayout + + return DeepseekV4CpuCacheLayout.load_from_args() if is_linear_att_mixed_model(get_env_start_args().model_dir): return LinearAttCacheConfig.load_from_args() raise ValueError("No hybrid state-cache layout registered for this model") diff --git a/lightllm/common/state_cache_manager/deepseek4.py b/lightllm/common/state_cache_manager/deepseek4.py new file mode 100644 index 0000000000..8d00848bf7 --- /dev/null +++ b/lightllm/common/state_cache_manager/deepseek4.py @@ -0,0 +1,30 @@ +import torch + +from .base import StateCacheManager + + +class DeepseekV4StateCacheManager(StateCacheManager): + """Pinned continuation checkpoints; compressed history stays in token pages.""" + + def __init__(self, size, layout, keep_num=0): + super().__init__(size, keep_num) + self.layout = layout + self.buffer = torch.empty((size, layout.page_nbytes - layout.swa_offset), dtype=torch.uint8, pin_memory=True) + + def get_state_cache(self, buffer_idx): + layout = self.layout + data = self.buffer[buffer_idx] + swa = data[: layout.swa_nbytes].view(layout.layer_num, layout.swa_gpu_pages_per_page, -1) + c4_start = layout.c4_state_offset - layout.swa_offset + indexer_start = layout.c4_indexer_state_offset - layout.swa_offset + c4 = ( + data[c4_start:indexer_start] + .view(torch.float32) + .view(layout.n_c4, layout.c4_state_rows, 4 * layout.head_dim) + ) + indexer = ( + data[indexer_start:] + .view(torch.float32) + .view(layout.n_c4, layout.c4_state_rows, 4 * layout.indexer_head_dim) + ) + return swa, c4, indexer diff --git a/lightllm/models/deepseek_v4/infer_struct.py b/lightllm/models/deepseek_v4/infer_struct.py index 9f4a13c5e2..a7e6be308d 100644 --- a/lightllm/models/deepseek_v4/infer_struct.py +++ b/lightllm/models/deepseek_v4/infer_struct.py @@ -94,13 +94,14 @@ def init_some_extra_state(self, model): image_left=self.dsv4_image_left, image_right=self.dsv4_image_right, ) + self.dsv4_swa_write_slots = torch.empty_like(pos, dtype=torch.int32) self.dsv4_swa_indices, self.dsv4_swa_lengths = build_swa_index( req_idx=self.dsv4_sparse_req_idx, positions=self.position_ids, - req_to_token_indexs=self.req_manager.req_to_token_indexs, - full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, + req_to_swa_pages=self.req_manager.req_to_swa_pages, swa_index=self.dsv4_swa_indices, swa_length=self.dsv4_swa_lengths, + swa_write_slots=self.dsv4_swa_write_slots, window=workspace.sliding_window, image_left=self.dsv4_image_left, image_right=self.dsv4_image_right, diff --git a/lightllm/models/deepseek_v4/layer_infer/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py index e2555a5283..26d64ee65c 100644 --- a/lightllm/models/deepseek_v4/layer_infer/compressor.py +++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py @@ -43,8 +43,7 @@ def _save_partial_states_kernel( token_to_batch_idx, b_req_idx, b_seq_len, - mem_index, - full_to_swa, + swa_write_slots, state_buffer, STATE_WIDTH: tl.constexpr, STATE_LAST_DIM: tl.constexpr, @@ -69,8 +68,7 @@ def _save_partial_states_kernel( return if IS_C4: - full_slot = tl.load(mem_index + token_idx).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + swa_slot = tl.load(swa_write_slots + token_idx).to(tl.int64) if swa_slot < 0: return state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) @@ -101,9 +99,8 @@ def _fused_compress_norm_rope_insert_kernel( b_seq_len, b_ready_cache_len, b_q_start_loc, - req_to_token, - req_to_token_stride0, - full_to_swa, + req_to_swa_pages, + req_to_swa_stride0, out_slots, norm_weight, rms_eps, @@ -170,12 +167,12 @@ def _fused_compress_norm_rope_insert_kernel( cache_pos = valid_pos if IS_C4: - full_slot = tl.load( - req_to_token + req_idx * req_to_token_stride0 + gather_pos, + swa_page = tl.load( + req_to_swa_pages + req_idx * req_to_swa_stride0 + gather_pos // SWA_PAGE_SIZE, mask=cache_pos, - other=0, + other=-1, ).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot, mask=cache_pos, other=-1).to(tl.int64) + swa_slot = swa_page * SWA_PAGE_SIZE + gather_pos % SWA_PAGE_SIZE state_row = (swa_slot // SWA_PAGE_SIZE) * STATE_RING + (swa_slot % STATE_RING) state_valid = cache_pos & (swa_slot >= 0) head_offset = tl.where(token_offsets >= COMPRESS_RATIO, HEAD_DIM, 0) @@ -362,7 +359,7 @@ def fused_compress( infer_state.b_ready_cache_len if infer_state.b_ready_cache_len is not None else infer_state.b_seq_len ) q_start_loc = infer_state.b_q_start_loc if infer_state.b_q_start_loc is not None else infer_state.b_seq_len - req_to_token_indexs = infer_state.req_manager.req_to_token_indexs + req_to_swa_pages = infer_state.req_manager.req_to_swa_pages _fused_compress_norm_rope_insert_kernel[(kv_score.shape[0],)]( kv_score, @@ -376,9 +373,8 @@ def fused_compress( infer_state.b_seq_len, ready_cache_len, q_start_loc, - req_to_token_indexs, - req_to_token_indexs.stride(0), - mem_manager.full_to_swa_indexs, + req_to_swa_pages, + req_to_swa_pages.stride(0), out_slots, norm_weight, eps, @@ -419,8 +415,7 @@ def fused_compress( token_to_batch_idx, infer_state.b_req_idx, infer_state.b_seq_len, - infer_state.mem_index, - mem_manager.full_to_swa_indexs, + infer_state.dsv4_swa_write_slots, state_buffer, STATE_WIDTH=state_width, STATE_LAST_DIM=state_last_dim, diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 579bebbfea..8c547c35bd 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -366,8 +366,7 @@ def _get_qkv( # kv: rmsnorm + rope + fp8 pack + scatter 进 swa 池,一个 DSV4 CUDA kernel 完成, infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, - mem_index=infer_state.mem_index, - swa_slots=getattr(infer_state, "dsv4_swa_write_slots", None), + swa_slots=infer_state.dsv4_swa_write_slots, kv=qkv[:, -self.head_dim_ :], kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, @@ -395,7 +394,7 @@ def context_attention_forward( self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): # _get_qkv writes the chunk's packed latent into the swa pool (fused kernel) before - # attention reads it back via full_to_swa indices (this custom forward bypasses the + # attention reads it back via request-owned SWA indices (this custom forward bypasses the # tpl _post_cache_kv path). q, q_lora, full_x = self._get_qkv(x, infer_state, layer_weight) o = self._context_attention_wrapper_run(q, q_lora, full_x, infer_state, layer_weight) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 4e58b55b7b..033851078b 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -293,7 +293,7 @@ def prepare_mtp_layer_hidden(self, layer_index: int, hidden): streams = hidden.view(-1, self.config["hc_mult"], self.config["hidden_size"]) return streams.mean(dim=1) - def _prepare_dsv4_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor) -> None: + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: if model_input.batch_size == 0 or (model_input.is_prefill and self.is_mtp_draft_model): return # Runtime inputs retain CPU mirrors. Synthetic warmup inputs are created on GPU. @@ -302,21 +302,22 @@ def _prepare_dsv4_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor if req_ids is None: req_ids = model_input.b_req_idx.cpu() seq_lens = model_input.b_seq_len.cpu() + if req_ids.numel() == 0: + return if model_input.is_prefill: ready_lens = model_input.b_ready_cache_len_cpu if ready_lens is None: ready_lens = model_input.b_ready_cache_len.cpu() - token_num = int((seq_lens - ready_lens).sum()) - self.req_manager.prepare_prefill(req_ids, ready_lens, seq_lens, mem_indexes[:token_num]) + self.req_manager.prepare_prefill(req_ids, ready_lens, seq_lens) else: mtp_indices = model_input.b_mtp_index_cpu if mtp_indices is None: mtp_indices = model_input.b_mtp_index.cpu() - self.req_manager.prepare_decode(req_ids, seq_lens, mtp_indices, mem_indexes[: len(req_ids)]) + self.req_manager.prepare_decode(req_ids, seq_lens, mtp_indices) def _select_mem_indexes(self, model_input: ModelInput): mem_indexes = super()._select_mem_indexes(model_input) - self._prepare_dsv4_slots(model_input, mem_indexes) + self._prepare_dsv4_slots(model_input) return mem_indexes def _init_to_get_rotary(self): diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py index 9b841159ac..956c4a2ab5 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py @@ -7,16 +7,14 @@ def _build_dspark_swa_index_kernel( req_idx_ptr, pos_ptr, - req_to_token_ptr, - req_to_token_stride0, - full_to_swa_ptr, + req_to_swa_pages, + req_to_swa_stride0, scratch_pages_ptr, swa_index_ptr, swa_index_stride0, swa_length_ptr, swa_write_slot_ptr, HOLD_REQ_ID: tl.constexpr, - HOLD_FULL_SLOT: tl.constexpr, HOLD_SWA_SLOT: tl.constexpr, WINDOW: tl.constexpr, BLOCK_SIZE: tl.constexpr, @@ -42,16 +40,12 @@ def _build_dspark_swa_index_kernel( source_position = tl.where(is_history, history_position, block_position) source_position = tl.where(is_hold | ~(is_history | is_block), 0, source_position) - history_full_slot = tl.load( - req_to_token_ptr + req_idx * req_to_token_stride0 + source_position, + history_page = tl.load( + req_to_swa_pages + req_idx * req_to_swa_stride0 + source_position // PAGE_SIZE, mask=in_width & is_history & ~is_hold, - other=HOLD_FULL_SLOT, - ).to(tl.int64) - history_swa_slot = tl.load( - full_to_swa_ptr + history_full_slot, - mask=in_width & is_history & ~is_hold, - other=HOLD_SWA_SLOT, + other=0, ) + history_swa_slot = history_page * PAGE_SIZE + source_position % PAGE_SIZE scratch_page = tl.load( scratch_pages_ptr + token_idx // BLOCK_SIZE, mask=~is_hold, @@ -74,8 +68,7 @@ def _build_dspark_swa_index_kernel( def build_dspark_swa_index( req_idx: torch.Tensor, positions: torch.Tensor, - req_to_token_indexs: torch.Tensor, - full_to_swa_indexs: torch.Tensor, + req_to_swa_pages: torch.Tensor, scratch_pages: torch.Tensor, swa_index: torch.Tensor, swa_length: torch.Tensor, @@ -84,7 +77,6 @@ def build_dspark_swa_index( block_size: int, page_size: int, hold_req_id: int, - hold_full_slot: int, hold_swa_slot: int, ): """Build ``history SWA + complete draft block`` indices for every DSpark query row.""" @@ -100,16 +92,14 @@ def build_dspark_swa_index( _build_dspark_swa_index_kernel[(token_num,)]( req_idx, positions, - req_to_token_indexs, - req_to_token_indexs.stride(0), - full_to_swa_indexs, + req_to_swa_pages, + req_to_swa_pages.stride(0), scratch_pages, swa_index, swa_index.stride(0), swa_length, swa_write_slots, HOLD_REQ_ID=hold_req_id, - HOLD_FULL_SLOT=hold_full_slot, HOLD_SWA_SLOT=hold_swa_slot, WINDOW=window, BLOCK_SIZE=block_size, diff --git a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py index a3ca231b0e..9d5ce47050 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py +++ b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py @@ -55,14 +55,15 @@ def build_image_visibility( def _build_swa_index_kernel( req_idx_ptr, pos_ptr, - req_to_token_ptr, - req_to_token_stride0, - full_to_swa_ptr, + req_to_swa_pages, + req_to_swa_stride0, image_left_ptr, image_right_ptr, swa_index_ptr, swa_index_stride0, swa_length_ptr, + swa_write_slots, + PAGE_SIZE: tl.constexpr, WINDOW: tl.constexpr, WIDTH: tl.constexpr, HAS_IMAGE: tl.constexpr, @@ -87,34 +88,28 @@ def _build_swa_index_kernel( valid = tl.where(is_image, w < image_length, valid) length = tl.where(is_image, image_length, length) safe_offset = tl.where(valid, offset, 0) - full_slot = tl.load(req_to_token_ptr + req * req_to_token_stride0 + safe_offset, mask=valid, other=0).to(tl.int64) - swa_slot = tl.load(full_to_swa_ptr + full_slot, mask=valid, other=-1) + page = tl.load(req_to_swa_pages + req * req_to_swa_stride0 + safe_offset // PAGE_SIZE, mask=valid, other=-1) + swa_slot = page * PAGE_SIZE + safe_offset % PAGE_SIZE out = tl.where(valid, swa_slot, -1).to(tl.int32) tl.store(swa_index_ptr + token_idx * swa_index_stride0 + w, out, mask=w_mask) tl.store(swa_length_ptr + token_idx, length.to(tl.int32)) + write_page = tl.load(req_to_swa_pages + req * req_to_swa_stride0 + pos // PAGE_SIZE) + tl.store(swa_write_slots + token_idx, write_page * PAGE_SIZE + pos % PAGE_SIZE) def build_swa_index( req_idx: torch.Tensor, positions: torch.Tensor, - req_to_token_indexs: torch.Tensor, - full_to_swa_indexs: torch.Tensor, + req_to_swa_pages: torch.Tensor, swa_index: torch.Tensor, swa_length: torch.Tensor, + swa_write_slots: torch.Tensor, window: int = None, image_left: torch.Tensor = None, image_right: torch.Tensor = None, ): - """Per-token sliding-window FlashMLA index table, built ONCE per forward (layer-independent: - full_to_swa is a single global map and the window is a model constant, so every layer's swa - indices are identical). Replaces DeepseekV4IndexInfer._swa_indices: for token t at - (req_idx, position) gather the last `window` tokens' full slots via req_to_token, then map - full -> swa; out-of-range positions store -1. - - Writes (swa_index [T, window] int32, swa_length [T] int32). The caller owns the output storage; - the reader adds the s_q axis via unsqueeze(1). - """ + """Build layer-independent SWA read/write slots from the request page table.""" T = positions.shape[0] width = swa_index.shape[1] window = width if window is None else int(window) @@ -127,14 +122,15 @@ def build_swa_index( _build_swa_index_kernel[(T,)]( req_idx, positions, - req_to_token_indexs, - req_to_token_indexs.stride(0), - full_to_swa_indexs, + req_to_swa_pages, + req_to_swa_pages.stride(0), image_left_arg, image_right_arg, swa_index, swa_index.stride(0), swa_length, + swa_write_slots, + PAGE_SIZE=128, WINDOW=window, WIDTH=width, HAS_IMAGE=image_left is not None, diff --git a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py index 6955a35643..c7c5996277 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py @@ -45,7 +45,9 @@ def _pack_gpu_cache_to_staging_kernel( c128_pool, c128_pool_stride0, c128_pool_stride1, - full_to_swa, + req_to_swa_pages, + req_to_swa_stride0, + source_req_meta, swa_pool, swa_pool_stride0, swa_pool_stride1, @@ -191,15 +193,12 @@ def _pack_gpu_cache_to_staging_kernel( layer = job // swa_gpu_page_num layer_i64 = layer.to(tl.int64) gpu_page_i64 = gpu_page.to(tl.int64) - full_slot = tl.load( - full_slots - + logical_page_i64 * token_page_size - + token_page_size - - history_block_size - + gpu_page_i64 * swa_pool_page_size - ).to(tl.int64) - pool_slot = tl.load(full_to_swa + full_slot).to(tl.int64) - physical_page = pool_slot // swa_pool_page_size + req_idx = tl.load(source_req_meta + logical_page_i64 * 2).to(tl.int64) + checkpoint_len = tl.load(source_req_meta + logical_page_i64 * 2 + 1).to(tl.int64) + position = checkpoint_len - history_block_size + gpu_page_i64 * swa_pool_page_size + physical_page = tl.load(req_to_swa_pages + req_idx * req_to_swa_stride0 + position // swa_pool_page_size).to( + tl.int64 + ) offsets = byte_block * BLOCK + offsets_base offsets_i64 = offsets.to(tl.int64) mask = offsets < swa_gpu_page_nbytes @@ -223,10 +222,12 @@ def _pack_gpu_cache_to_staging_kernel( layer = job // 4 row_i64 = row.to(tl.int64) layer_i64 = layer.to(tl.int64) - full_slot = tl.load(full_slots + logical_page_i64 * token_page_size + token_page_size - 4 + row_i64).to( + req_idx = tl.load(source_req_meta + logical_page_i64 * 2).to(tl.int64) + position = tl.load(source_req_meta + logical_page_i64 * 2 + 1).to(tl.int64) - 4 + row_i64 + page = tl.load(req_to_swa_pages + req_idx * req_to_swa_stride0 + position // swa_pool_page_size).to( tl.int64 ) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + swa_slot = page * swa_pool_page_size + position % swa_pool_page_size state_row = (swa_slot // swa_pool_page_size) * c4_state_ring + swa_slot % c4_state_ring offsets = byte_block * BLOCK + offsets_base offsets_i64 = offsets.to(tl.int64) @@ -479,13 +480,16 @@ def _unpack_cpu_cache_to_gpu_kernel( tl.store(indexer_target, tl.load(indexer_source, mask=indexer_mask), mask=indexer_mask) -def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, staging: torch.Tensor) -> None: +def pack_gpu_cache_to_staging( + mem_manager, source_mem_indexes: torch.Tensor, source_req_meta: torch.Tensor, staging: torch.Tensor +) -> None: """Pack a compact batch of complete CPU checkpoint pages into caller-owned CUDA staging.""" layout = mem_manager.cpu_cache_layout assert source_mem_indexes.is_cuda and staging.is_cuda assert source_mem_indexes.ndim == 2 and source_mem_indexes.shape[1] == layout.token_page_size assert staging.dtype == torch.uint8 and staging.is_contiguous() page_num = source_mem_indexes.shape[0] + assert source_req_meta.shape == (page_num, 2) and source_req_meta.is_cuda assert staging.shape == (page_num, layout.page_nbytes) full_slots = source_mem_indexes.reshape(-1) @@ -522,7 +526,9 @@ def pack_gpu_cache_to_staging(mem_manager, source_mem_indexes: torch.Tensor, sta c128_pool, c128_pool.stride(0) if has_c128 else 0, c128_pool.stride(1) if has_c128 else 0, - mem_manager.full_to_swa_indexs, + mem_manager.req_to_swa_pages, + mem_manager.req_to_swa_pages.stride(0), + source_req_meta, swa_pool, swa_pool.stride(0), swa_pool.stride(1), diff --git a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py index 5df62f4131..cb35e8ea48 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py @@ -10,8 +10,9 @@ _BYTE_BLOCK = 8192 # Per-source row: c4 data/indexer, c128 data, SWA map/data, c4 state/indexer state. _SOURCE_POOL_PTR_COUNT = 7 -# Per-task row: source manager index, token count, source full-slot pointer, destination full-slot pointer. -_TASK_META_WIDTH = 4 +# Per-task row: source manager index, source/destination full-slot pointers, +# source/destination request IDs, logical end. +_TASK_META_WIDTH = 6 @triton.jit @@ -28,7 +29,8 @@ def _copy_dsv4_dp_caches_kernel( dst_c128_pool, dst_c128_pool_stride0, dst_c128_pool_stride1, - dst_full_to_swa, + dst_req_to_swa, + req_swa_stride0, dst_swa_pool, dst_swa_pool_stride0, dst_swa_pool_stride1, @@ -79,8 +81,8 @@ def _copy_dsv4_dp_caches_kernel( history_block = tl.load(history_meta + history_index * 2 + 1).to(tl.int64) task_row = task_meta + task * task_meta_width source_manager = tl.load(task_row).to(tl.int64) - src_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) - dst_full_slots = tl.load(task_row + 3).to(tl.pointer_type(tl.int32)) + src_full_slots = tl.load(task_row + 1).to(tl.pointer_type(tl.int32)) + dst_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count layer_i64 = layer.to(tl.int64) @@ -161,11 +163,11 @@ def _copy_dsv4_dp_caches_kernel( task_pid = tail_pid // task_num task_row = task_meta + task * task_meta_width source_manager = tl.load(task_row).to(tl.int64) - token_num = tl.load(task_row + 1).to(tl.int64) - src_full_slots = tl.load(task_row + 2).to(tl.pointer_type(tl.int32)) - dst_full_slots = tl.load(task_row + 3).to(tl.pointer_type(tl.int32)) source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count - src_full_to_swa = tl.load(source_ptr_row + 3).to(tl.pointer_type(tl.int32)) + src_req_to_swa = tl.load(source_ptr_row + 3).to(tl.pointer_type(tl.int32)) + src_req = tl.load(task_row + 3).to(tl.int64) + dst_req = tl.load(task_row + 4).to(tl.int64) + end = tl.load(task_row + 5).to(tl.int64) src_swa_pool = tl.load(source_ptr_row + 4).to(tl.pointer_type(tl.uint8)) if task_pid < swa_program_num: @@ -173,13 +175,9 @@ def _copy_dsv4_dp_caches_kernel( layer = task_pid // 2 page_i64 = page.to(tl.int64) layer_i64 = layer.to(tl.int64) - full_offset = token_num - history_block_size + page_i64 * swa_pool_page_size - src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) - dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_pool_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) - dst_pool_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) - src_page = src_pool_slot // swa_pool_page_size - dst_page = dst_pool_slot // swa_pool_page_size + position = end - history_block_size + page_i64 * swa_pool_page_size + src_page = tl.load(src_req_to_swa + src_req * req_swa_stride0 + position // swa_pool_page_size).to(tl.int64) + dst_page = tl.load(dst_req_to_swa + dst_req * req_swa_stride0 + position // swa_pool_page_size).to(tl.int64) src_page_ptr = src_swa_pool + layer_i64 * dst_swa_pool_stride0 + src_page * dst_swa_pool_stride1 dst_page_ptr = dst_swa_pool + layer_i64 * dst_swa_pool_stride0 + dst_page * dst_swa_pool_stride1 for byte_start in tl.range(0, swa_pool_page_nbytes, BLOCK): @@ -197,11 +195,15 @@ def _copy_dsv4_dp_caches_kernel( layer_i64 = layer.to(tl.int64) src_c4_state = tl.load(source_ptr_row + 5).to(tl.pointer_type(tl.uint8)) src_c4_indexer_state = tl.load(source_ptr_row + 6).to(tl.pointer_type(tl.uint8)) - full_offset = token_num - 4 + row_i64 - src_full_slot = tl.load(src_full_slots + full_offset).to(tl.int64) - dst_full_slot = tl.load(dst_full_slots + full_offset).to(tl.int64) - src_swa_slot = tl.load(src_full_to_swa + src_full_slot).to(tl.int64) - dst_swa_slot = tl.load(dst_full_to_swa + dst_full_slot).to(tl.int64) + position = end - 4 + row_i64 + src_page = tl.load(src_req_to_swa + src_req * req_swa_stride0 + position // swa_pool_page_size).to( + tl.int64 + ) + dst_page = tl.load(dst_req_to_swa + dst_req * req_swa_stride0 + position // swa_pool_page_size).to( + tl.int64 + ) + src_swa_slot = src_page * swa_pool_page_size + position % swa_pool_page_size + dst_swa_slot = dst_page * swa_pool_page_size + position % swa_pool_page_size src_state_row = (src_swa_slot // swa_pool_page_size) * c4_state_ring + src_swa_slot % c4_state_ring dst_state_row = (dst_swa_slot // swa_pool_page_size) * c4_state_ring + dst_swa_slot % c4_state_ring @@ -279,7 +281,8 @@ def copy_dsv4_dp_caches( dst_c128_pool, dst_c128_pool.stride(0) if has_c128 else 0, dst_c128_pool.stride(1) if has_c128 else 0, - dst_mem_manager.full_to_swa_indexs, + dst_mem_manager.req_to_swa_pages, + dst_mem_manager.req_to_swa_pages.stride(0), dst_swa_pool, dst_swa_pool.stride(0), dst_swa_pool.stride(1), diff --git a/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py index 33e5eb5e67..a025bcf1c1 100644 --- a/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py +++ b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py @@ -17,8 +17,9 @@ @triton.jit def _pd_tail_kernel( - full_slots, - full_to_swa, + req_to_swa_pages, + req_to_swa_stride0, + start_kv_index, swa_pool, swa_pool_stride0, swa_pool_stride1, @@ -39,8 +40,7 @@ def _pd_tail_kernel( swa_first_full_offset, swa_section_gpu_page_start, c4_row_num, - c4_first_full_offset, - c4_section_row_start, + c4_rows, c128_row_num, c128_first_position, c128_section_row_start, @@ -77,9 +77,10 @@ def _pd_tail_kernel( layer = job // swa_page_num gpu_page_i64 = gpu_page.to(tl.int64) layer_i64 = layer.to(tl.int64) - full_slot = tl.load(full_slots + swa_first_full_offset + gpu_page_i64 * swa_page_size).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) - physical_page = swa_slot // swa_page_size + position = start_kv_index + swa_first_full_offset + gpu_page_i64 * swa_page_size + physical_page = tl.load(req_to_swa_pages + req_idx * req_to_swa_stride0 + position // swa_page_size).to( + tl.int64 + ) offsets = byte_block * BYTE_BLOCK + tl.arange(0, BYTE_BLOCK) offsets_i64 = offsets.to(tl.int64) mask = offsets < swa_page_nbytes @@ -105,8 +106,10 @@ def _pd_tail_kernel( layer = job // c4_row_num row_i64 = row.to(tl.int64) layer_i64 = layer.to(tl.int64) - full_slot = tl.load(full_slots + c4_first_full_offset + row_i64).to(tl.int64) - swa_slot = tl.load(full_to_swa + full_slot).to(tl.int64) + position = tl.load(c4_rows + row_i64 * 2).to(tl.int64) + section_row = tl.load(c4_rows + row_i64 * 2 + 1).to(tl.int64) + page = tl.load(req_to_swa_pages + req_idx * req_to_swa_stride0 + position // swa_page_size).to(tl.int64) + swa_slot = page * swa_page_size + position % swa_page_size state_row = (swa_slot // swa_page_size) * c4_state_ring + swa_slot % c4_state_ring offsets = state_block * STATE_BLOCK + tl.arange(0, STATE_BLOCK) offsets_i64 = offsets.to(tl.int64) @@ -116,7 +119,7 @@ def _pd_tail_kernel( staging_f32 + c4_state_section_offset_f32 + layer_i64 * c4_state_section_layer_elems - + (c4_section_row_start + row_i64) * c4_state_width + + section_row * c4_state_width + offsets_i64 ) indexer_mask = offsets < c4_indexer_state_width @@ -130,7 +133,7 @@ def _pd_tail_kernel( staging_f32 + c4_indexer_state_section_offset_f32 + layer_i64 * c4_indexer_state_section_layer_elems - + (c4_section_row_start + row_i64) * c4_indexer_state_width + + section_row * c4_indexer_state_width + offsets_i64 ) if MODE == 0: @@ -179,7 +182,6 @@ def _copy_pd_tail( mode, mem_manager, layout, - full_slots, staging, start_kv_index, end_kv_index, @@ -196,20 +198,25 @@ def _copy_pd_tail( c4_state = None c4_indexer_state = None c4_row_num = 0 - c4_first_full_offset = 0 - c4_section_row_start = 0 + c4_rows = None if mem_manager.c4_pool is not None: + checkpoint_len = (request_kv_len - 1) // _C4_TOKEN_BLOCK * _C4_TOKEN_BLOCK + checkpoint_positions = list(range(checkpoint_len - 4, checkpoint_len)) if checkpoint_len else [] c4_remainder = request_kv_len % _C4_RATIO - c4_required_rows = _C4_RATIO if c4_remainder == 0 else _C4_RATIO + c4_remainder + c4_required_rows = _C4_RATIO + c4_remainder c4_state_start = max(0, request_kv_len - c4_required_rows) - c4_intersection_start = max(start_kv_index, c4_state_start) - c4_intersection_end = min(end_kv_index, request_kv_len) - if c4_intersection_start < c4_intersection_end: + rows = [(position, row) for row, position in enumerate(checkpoint_positions)] + rows.extend( + (position, 4 + row) + for row, position in enumerate(range(c4_state_start, request_kv_len)) + if position not in checkpoint_positions + ) + rows = [(position, row) for position, row in rows if start_kv_index <= position < end_kv_index] + if rows: c4_state = mem_manager.c4_state_buffer c4_indexer_state = mem_manager.c4_indexer_state_buffer - c4_row_num = c4_intersection_end - c4_intersection_start - c4_first_full_offset = c4_intersection_start - start_kv_index - c4_section_row_start = c4_intersection_start - c4_state_start + c4_row_num = len(rows) + c4_rows = torch.tensor(rows, dtype=torch.int64, device="cuda") c128_state = None c128_row_num = 0 @@ -242,8 +249,9 @@ def _copy_pd_tail( c128_program_num = c128_state.shape[0] * c128_row_num * c128_blocks_per_row if c128_state is not None else 0 _pd_tail_kernel[(swa_program_num + c4_program_num + c128_program_num,)]( - full_slots, - mem_manager.full_to_swa_indexs, + mem_manager.req_to_swa_pages, + mem_manager.req_to_swa_pages.stride(0), + start_kv_index, swa_pool, swa_pool.stride(0), swa_pool.stride(1), @@ -264,8 +272,7 @@ def _copy_pd_tail( swa_page_start - start_kv_index, (swa_page_start - swa_tail_start) // _SWA_PAGE_SIZE, c4_row_num, - c4_first_full_offset, - c4_section_row_start, + c4_rows, c128_row_num, c128_first_position, c128_section_row_start, @@ -346,7 +353,7 @@ def _copy_pd_cache_page( c128_section_layer_nbytes=layout.c128_layer_nbytes, ) - swa_tail_start = max(0, request_kv_len // _C4_TOKEN_BLOCK * _C4_TOKEN_BLOCK - _C4_TOKEN_BLOCK) + swa_tail_start = max(0, (request_kv_len - 1) // _C4_TOKEN_BLOCK * _C4_TOKEN_BLOCK - _C4_TOKEN_BLOCK) if end_kv_index <= swa_tail_start: return layout.swa_offset @@ -354,7 +361,6 @@ def _copy_pd_cache_page( mode, mem_manager, layout, - full_slots, staging, start_kv_index, end_kv_index, diff --git a/lightllm/models/deepseek_v4_dspark/infer_struct.py b/lightllm/models/deepseek_v4_dspark/infer_struct.py index 175fc46163..0367243fbf 100644 --- a/lightllm/models/deepseek_v4_dspark/infer_struct.py +++ b/lightllm/models/deepseek_v4_dspark/infer_struct.py @@ -28,8 +28,7 @@ def init_some_extra_state(self, model): build_dspark_swa_index( req_idx=self.dsv4_sparse_req_idx, positions=self.position_ids, - req_to_token_indexs=self.req_manager.req_to_token_indexs, - full_to_swa_indexs=self.mem_manager.full_to_swa_indexs, + req_to_swa_pages=self.req_manager.req_to_swa_pages, scratch_pages=self.mtp_draft_swa_pages, swa_index=self.dsv4_swa_indices, swa_length=self.dsv4_swa_lengths, @@ -38,6 +37,5 @@ def init_some_extra_state(self, model): block_size=model.block_size, page_size=DSV4_SWA_PAGE_SIZE, hold_req_id=self.req_manager.HOLD_REQUEST_ID, - hold_full_slot=self.mem_manager.HOLD_TOKEN_MEMINDEX, hold_swa_slot=self.mem_manager.swa_pool.HOLD_TOKEN_MEMINDEX, ) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 7ac58c364d..40866be713 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -183,7 +183,7 @@ def _init_custom(self): layer.sin_compress_table = self._sin_cached_compress self.layers_infer[0].context_wkv_weight = self.context_wkv_weight - def _prepare_dsv4_slots(self, model_input: ModelInput, mem_indexes: torch.Tensor) -> None: + def _prepare_dsv4_slots(self, model_input: ModelInput) -> None: # Target-hidden commits use target SWA; proposal blocks own separate scratch pages. return diff --git a/lightllm/models/deepseek_v4_mtp/model.py b/lightllm/models/deepseek_v4_mtp/model.py index bae8085497..2dc54e71fa 100644 --- a/lightllm/models/deepseek_v4_mtp/model.py +++ b/lightllm/models/deepseek_v4_mtp/model.py @@ -7,7 +7,6 @@ from tqdm import tqdm from lightllm.common.basemodel import TpPartBaseModel -from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel from lightllm.models.deepseek_v4_mtp.layer_infer.pre_layer_infer import DeepseekV4MTPPreLayerInfer from lightllm.models.deepseek_v4_mtp.layer_infer.transformer_layer_infer import ( @@ -101,11 +100,6 @@ def _init_some_value(self): self.layers_num = 1 return - def _microbatch_overlap_decode_cuda(self, model_input0: ModelInput, model_input1: ModelInput): - self._prepare_dsv4_slots(model_input0) - self._prepare_dsv4_slots(model_input1) - return super()._microbatch_overlap_decode_cuda(model_input0, model_input1) - def _gen_special_model_input(self, token_num: int): return { "mtp_draft_input_hiddens": torch.randn( diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index c324798e8b..abcb074a7d 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -28,7 +28,6 @@ ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args from lightllm.utils.device_utils import is_sm100_gpu, is_sm90_gpu -from lightllm.utils.auto_shm_cleanup import register_sysv_shm_for_cleanup logger = init_logger(__name__) @@ -101,7 +100,6 @@ def _launch_subprocesses(args: StartArgs): if args.enable_cpu_cache: # 生成一个用于创建cpu kv cache的共享内存id。 args.cpu_kv_cache_shm_id = uuid.uuid1().int % 123456789 - register_sysv_shm_for_cleanup(args.cpu_kv_cache_shm_id) if args.enable_multimodal: args.multi_modal_cache_shm_id = uuid.uuid1().int % 123456789 @@ -217,8 +215,9 @@ def _launch_subprocesses(args: StartArgs): if args.page_size < 1: raise ValueError(f"--page_size must be >= 1, got {args.page_size}") - if get_model_type(args.model_dir) == "deepseek_v4" and args.page_size != 256: - raise ValueError("DeepSeek-V4 requires --page_size 256") + if get_model_type(args.model_dir) == "deepseek_v4": + if args.page_size != 256 or args.linear_att_hash_page_size != 256: + raise ValueError("DeepSeek-V4 requires --page_size 256 --linear_att_hash_page_size 256") if args.run_mode in ("prefill", "decode"): assert args.pd_kv_page_size % args.page_size == 0, "--pd_kv_page_size must be divisible by --page_size" @@ -337,9 +336,22 @@ def _launch_subprocesses(args: StartArgs): # 避免请求释放时将不完整的大页 state 写入 radix cache 并触发断言。 args.linear_att_page_block_num = 10000000 - if args.enable_cpu_cache and is_hybrid_att_model(args.model_dir): + if ( + args.enable_cpu_cache + and is_hybrid_att_model(args.model_dir) + and get_model_type(args.model_dir) != "deepseek_v4" + ): args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for hybrid att model") + elif args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4": + big_page_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + if big_page_tokens <= args.max_req_total_len: + if args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = big_page_tokens + if args.cpu_cache_token_page_size != big_page_tokens: + raise ValueError("DeepSeek-V4 CPU cache pages must match the hybrid big-page checkpoint interval") + elif args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = 2048 elif args.enable_cpu_cache and args.cpu_cache_token_page_size is None: args.cpu_cache_token_page_size = 2048 if get_model_type(args.model_dir) == "deepseek_v4" else 256 if args.enable_cpu_cache: diff --git a/lightllm/server/core/objs/req.py b/lightllm/server/core/objs/req.py index b3f954ec51..128c966007 100644 --- a/lightllm/server/core/objs/req.py +++ b/lightllm/server/core/objs/req.py @@ -10,7 +10,7 @@ from .token_chunck_hash_list import TokenHashList, CpuCachePageList, TokenPageLenList from lightllm.server.req_id_generator import convert_sub_id_to_group_id from lightllm.utils.envs_utils import get_env_start_args -from lightllm.utils.config_utils import is_hybrid_att_model +from lightllm.utils.config_utils import is_hybrid_att_model, get_model_type from lightllm.utils.kv_cache_utils import compute_token_list_hash from typing import Any, Dict, List, Union from lightllm.utils.log_utils import init_logger @@ -217,6 +217,7 @@ def init( args = get_env_start_args() if is_hybrid_att_model(args.model_dir): self._fill_hybrid_token_hash() + if is_hybrid_att_model(args.model_dir) and get_model_type(args.model_dir) != "deepseek_v4": if args.enable_cpu_cache: cpu_cache_hash_list, cpu_cache_page_len_list = self._calcu_hybrid_cpu_cache_page_len_list() self.token_hash_list = TokenHashList() diff --git a/lightllm/server/router/dynamic_prompt/radix_cache.py b/lightllm/server/router/dynamic_prompt/radix_cache.py index 40caf8ee10..b954912a9f 100644 --- a/lightllm/server/router/dynamic_prompt/radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/radix_cache.py @@ -2,12 +2,9 @@ import torch import numpy as np import collections -from typing import Any, Tuple, Dict, Set, List, Optional, Union +from typing import Tuple, Dict, Set, List, Optional, Union from sortedcontainers import SortedSet from .shared_arr import SharedArray -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) class UniqueTimeIdGenerator: @@ -31,7 +28,6 @@ def __init__(self, page_size: int = 1): self.parent: TreeNode = None self.token_id_key: torch.Tensor = None self.token_mem_index_value: torch.Tensor = None # 用于记录存储的 token_index 为每个元素在 token mem 中的index位置 - self.token_extra_value: Any = None self.ref_counter = 0 self.time_id = time_gen.generate_time_id() # 用于标识时间周期 @@ -47,16 +43,13 @@ def get_child_key(self, token_ids: torch.Tensor): return first_page.item() return first_page.numpy().tobytes() - def split_node(self, prefix_len, child_key_fn=None, extra_value_ops=None): + def split_node(self, prefix_len): assert prefix_len > 0 and prefix_len % self.page_size == 0 split_parent_node = TreeNode(page_size=self.page_size) split_parent_node.parent = self.parent split_parent_node.parent.children[self.get_child_key(self.token_id_key)] = split_parent_node split_parent_node.token_id_key = self.token_id_key[0:prefix_len] split_parent_node.token_mem_index_value = self.token_mem_index_value[0:prefix_len] - if self.token_extra_value is not None and extra_value_ops is not None: - split_parent_node.token_extra_value = extra_value_ops.slice(self.token_extra_value, 0, prefix_len) - self.token_extra_value = extra_value_ops.slice(self.token_extra_value, prefix_len, len(self.token_id_key)) split_parent_node.children = {} split_parent_node.children[self.get_child_key(self.token_id_key[prefix_len:])] = self split_parent_node.ref_counter = self.ref_counter @@ -73,11 +66,10 @@ def split_node(self, prefix_len, child_key_fn=None, extra_value_ops=None): self.node_prefix_total_len = self.parent.node_prefix_total_len + new_len return split_parent_node - def add_and_return_new_child(self, token_id_key, token_mem_index_value, token_extra_value=None, child_key=None): + def add_and_return_new_child(self, token_id_key, token_mem_index_value): child = TreeNode(page_size=self.page_size) child.token_id_key = token_id_key child.token_mem_index_value = token_mem_index_value - child.token_extra_value = token_extra_value child_key = child.get_child_key(child.token_id_key) assert child_key not in self.children.keys() self.children[child_key] = child @@ -117,7 +109,7 @@ def match(t1: torch.Tensor, t2: torch.Tensor) -> int: class RadixCache: - def __init__(self, total_token_num, rank_in_node, mem_manager=None, page_size: int = 1, extra_value_ops=None): + def __init__(self, total_token_num, rank_in_node, mem_manager=None, page_size: int = 1): from lightllm.common.kv_cache_mem_manager import MemoryManager self.total_token_num = total_token_num @@ -131,7 +123,6 @@ def __init__(self, total_token_num, rank_in_node, mem_manager=None, page_size: i f"RadixCache page_size {page_size} must match mem_manager page_size {mem_manager.page_size}" ) self.page_size = page_size - self.extra_value_ops = extra_value_ops self.root_node = TreeNode(page_size=page_size) self.root_node.token_id_key = torch.zeros((0,), device="cpu", dtype=self._key_dtype) @@ -145,99 +136,34 @@ def __init__(self, total_token_num, rank_in_node, mem_manager=None, page_size: i self.refed_tokens_num.arr[0] = 0 self.tree_total_tokens_num = SharedArray(f"tree_total_tokens_num_{rank_in_node}", (1,), dtype=np.int64) self.tree_total_tokens_num.arr[0] = 0 - self.swa_tree_total_pages_num = 0 - self.swa_refed_pages_num = 0 - # 每个 prompt-cache 页折算多少 swa 页(DSV4 为 256/128=2);非 swa 场景为 0,_node_swa_pages_num 退化为常数 0。 - self._swa_pages_per_prompt_page = self._probe_swa_pages_per_prompt_page() - - def _probe_swa_pages_per_prompt_page(self) -> int: - """构造期探测一次 mem_manager 是否带 swa_pool,缓存折算系数,避免热路径反复 getattr。""" - if self.mem_manager is None or self.extra_value_ops is None: - return 0 - swa_pool = getattr(self.mem_manager, "swa_pool", None) - swa_page_size = getattr(swa_pool, "page_size", None) - if swa_page_size is None: - return 0 - return (self.page_size + int(swa_page_size) - 1) // int(swa_page_size) - - def _node_swa_pages_num(self, node: TreeNode) -> int: - if self._swa_pages_per_prompt_page == 0 or node.token_extra_value is None: - return 0 - valid = node.token_extra_value.swa_page_valid - if valid is None: - return 0 - return int(valid.sum().item()) * self._swa_pages_per_prompt_page - - def _node_direct_ref_num(self, node: TreeNode) -> int: - # 子树引用同时计入父子节点,差值就是直接命中本节点的请求数。 - return node.ref_counter - sum(child.ref_counter for child in node.children.values()) - - def _node_last_valid_swa_pages_num(self, node: TreeNode) -> int: - if node.token_extra_value is None: - return 0 - return self._swa_pages_per_prompt_page if node.token_extra_value.swa_last_valid_page >= 0 else 0 - - def _align_len(self, length: int) -> int: - if self.page_size <= 1: - return int(length) - return int(length) // self.page_size * self.page_size - - def align_len(self, length: int) -> int: - return self._align_len(length) - - def _child_key(self, key: torch.Tensor): - if self.page_size <= 1: - return key[0].item() - return key[: self.page_size].numpy().tobytes() - - def _match_len(self, key: torch.Tensor, node_key: torch.Tensor) -> int: - prefix_len = match(key, node_key) - return self._align_len(prefix_len) - - def _slice_extra(self, extra_value, start: int, end: int): - if extra_value is None: - return None - assert self.extra_value_ops is not None - return self.extra_value_ops.slice(extra_value, start, end) - - def _concat_extra(self, values: list): - values = [v for v in values if v is not None] - if len(values) == 0: - return None - assert self.extra_value_ops is not None - return self.extra_value_ops.concat(values) - def insert(self, key, value=None, extra_value=None) -> Tuple[int, Optional[TreeNode]]: + def insert(self, key, value=None) -> Tuple[int, Optional[TreeNode]]: if value is None: value = key - align_len = self._align_len(len(key)) - key = key[:align_len] - value = value[:align_len] - if extra_value is not None: - extra_value = self._slice_extra(extra_value, 0, align_len) - assert len(key) == len(value) # and len(key) >= 1 aligned_len = len(key) // self.page_size * self.page_size if aligned_len == 0: return 0, None - return self._insert_helper(self.root_node, key, value, extra_value) + key = key[:aligned_len] + value = value[:aligned_len] + return self._insert_helper(self.root_node, key, value) - def _insert_helper(self, node: TreeNode, key, value, extra_value) -> Tuple[int, Optional[TreeNode]]: + def _insert_helper(self, node: TreeNode, key, value) -> Tuple[int, Optional[TreeNode]]: handle_stack = collections.deque() update_list = collections.deque() - handle_stack.append((node, key, value, extra_value)) + handle_stack.append((node, key, value)) ans_prefix_len = 0 ans_node = None while len(handle_stack) != 0: - node, key, value, extra_value = handle_stack.popleft() - ans_tuple = self._insert_helper_no_recursion(node=node, key=key, value=value, extra_value=extra_value) - if len(ans_tuple) == 5: - (_prefix_len, new_node, new_key, new_value, new_extra_value) = ans_tuple + node, key, value = handle_stack.popleft() + ans_tuple = self._insert_helper_no_recursion(node=node, key=key, value=value) + if len(ans_tuple) == 4: + (_prefix_len, new_node, new_key, new_value) = ans_tuple ans_prefix_len += _prefix_len - handle_stack.append((new_node, new_key, new_value, new_extra_value)) + handle_stack.append((new_node, new_key, new_value)) else: _prefix_len, ans_node = ans_tuple ans_prefix_len += _prefix_len @@ -255,8 +181,8 @@ def _insert_helper(self, node: TreeNode, key, value, extra_value) -> Tuple[int, return ans_prefix_len, ans_node def _insert_helper_no_recursion( - self, node: TreeNode, key: torch.Tensor, value: torch.Tensor, extra_value=None - ) -> Union[Tuple[int, Optional[TreeNode]], Tuple[int, TreeNode, torch.Tensor, torch.Tensor, Any]]: + self, node: TreeNode, key: torch.Tensor, value: torch.Tensor + ) -> Union[Tuple[int, Optional[TreeNode]], Tuple[int, TreeNode, torch.Tensor, torch.Tensor]]: if node.is_leaf(): self.evict_tree_set.discard(node) @@ -275,14 +201,10 @@ def _insert_helper_no_recursion( self.evict_tree_set.add(child) return prefix_len, child elif prefix_len < len(child.token_id_key): - if prefix_len == 0: - return 0, node if child.is_leaf(): self.evict_tree_set.discard(child) - split_parent_node = child.split_node( - prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops - ) + split_parent_node = child.split_node(prefix_len) if split_parent_node.is_leaf(): self.evict_tree_set.add(split_parent_node) @@ -294,26 +216,15 @@ def _insert_helper_no_recursion( assert False, "can not run to here" elif prefix_len < len(key) and prefix_len < len(child.token_id_key): - if prefix_len == 0: - return 0, node if child.is_leaf(): self.evict_tree_set.discard(child) - new_extra_value = self._slice_extra(extra_value, prefix_len, len(key)) key = key[prefix_len:] value = value[prefix_len:] - split_parent_node = child.split_node( - prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops - ) - new_node = split_parent_node.add_and_return_new_child( - key, - value, - token_extra_value=new_extra_value, - child_key=self._child_key(key), - ) + split_parent_node = child.split_node(prefix_len) + new_node = split_parent_node.add_and_return_new_child(key, value) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) - self.swa_tree_total_pages_num += self._node_swa_pages_num(new_node) if split_parent_node.is_leaf(): self.evict_tree_set.add(split_parent_node) @@ -324,42 +235,26 @@ def _insert_helper_no_recursion( self.evict_tree_set.add(child) return prefix_len, new_node elif prefix_len < len(key) and prefix_len == len(child.token_id_key): - return ( - prefix_len, - child, - key[prefix_len:], - value[prefix_len:], - self._slice_extra(extra_value, prefix_len, len(key)), - ) + return (prefix_len, child, key[prefix_len:], value[prefix_len:]) else: assert False, "can not run to here" else: - new_node = node.add_and_return_new_child( - key, - value, - token_extra_value=extra_value, - child_key=first_key_id, - ) + new_node = node.add_and_return_new_child(key, value) # update total token num self.tree_total_tokens_num.arr[0] += len(new_node.token_mem_index_value) - self.swa_tree_total_pages_num += self._node_swa_pages_num(new_node) if new_node.is_leaf(): self.evict_tree_set.add(new_node) return 0, new_node def match_prefix(self, key, update_refs=False): - key = key[: self._align_len(len(key))] - if len(key) == 0: - return None, 0, None - key = self._trim_key_by_extra_value_validity(key) - if len(key) == 0: + aligned_len = len(key) // self.page_size * self.page_size + if aligned_len == 0: return None, 0, None + key = key[:aligned_len] ans_value_list = [] tree_node = self._match_prefix_helper(self.root_node, key, ans_value_list, update_refs=update_refs) if tree_node != self.root_node: - if update_refs and self._swa_pages_per_prompt_page > 0 and self._node_direct_ref_num(tree_node) == 1: - self.swa_refed_pages_num += self._node_last_valid_swa_pages_num(tree_node) if len(ans_value_list) != 0: value = torch.concat(ans_value_list) else: @@ -370,30 +265,6 @@ def match_prefix(self, key, update_refs=False): self.dec_node_ref_counter(self.root_node) return None, 0, None - def _trim_key_by_extra_value_validity(self, key: torch.Tensor) -> torch.Tensor: - """命中有效性裁剪(extra_value_ops 提供 valid_match_length 时启用,如 DeepSeek-V4 的 - swa 按页 bitmap): 先做一次只读探测遍历得到自然命中与沿路 extra_value,按其有效边界截短 - key,随后的正常遍历(加引用/分裂)只走截短后的前缀 —— 引用计数与最终返回值在同一次遍历 - 内保持一致,不存在事后裁剪导致的漏减/多减。 - - 探测遍历可能分裂部分命中的节点(与正常遍历同语义,树不变式不受影响)。裁剪只会缩短命中, - 没有任何失败路径。""" - if self.extra_value_ops is None: - return key - valid_match_length = getattr(self.extra_value_ops, "valid_match_length", None) - if valid_match_length is None: - return key - probe_values = [] - probe_node = self._match_prefix_helper(self.root_node, key, probe_values, update_refs=False) - if probe_node == self.root_node or len(probe_values) == 0: - return key - natural_len = sum(len(v) for v in probe_values) - extra_value = self.get_extra_value_by_node(probe_node) - valid_len = int(valid_match_length(extra_value, natural_len)) - if valid_len < natural_len: - return key[:valid_len] - return key - def _match_prefix_helper( self, node: TreeNode, key: torch.Tensor, ans_value_list: list, update_refs=False ) -> TreeNode: @@ -451,14 +322,10 @@ def _match_prefix_helper_no_recursion( ans_value_list.append(child.token_mem_index_value) return (child, key[prefix_len:]) elif prefix_len < len(child.token_id_key): - if prefix_len == 0: - return node if child.is_leaf(): self.evict_tree_set.discard(child) - split_parent_node = child.split_node( - prefix_len, child_key_fn=self._child_key, extra_value_ops=self.extra_value_ops - ) + split_parent_node = child.split_node(prefix_len) ans_value_list.append(split_parent_node.token_mem_index_value) if update_refs: @@ -489,11 +356,8 @@ def evict(self, need_remove_tokens, evict_callback): ), "error evict tree node state" num_evicted += len(node.token_mem_index_value) evict_callback(node.token_mem_index_value) - if self.extra_value_ops is not None and node.token_extra_value is not None: - self.extra_value_ops.free(node.token_extra_value) # update total token num self.tree_total_tokens_num.arr[0] -= len(node.token_mem_index_value) - self.swa_tree_total_pages_num -= self._node_swa_pages_num(node) parent_node: TreeNode = node.parent parent_node.remove_child(node) if parent_node.is_leaf(): @@ -527,7 +391,6 @@ def _try_merge(self, child_node: TreeNode) -> Optional[TreeNode]: child_node.token_mem_index_value = torch.cat( [parent_node.token_mem_index_value, child_node.token_mem_index_value] ) - child_node.token_extra_value = self._concat_extra([parent_node.token_extra_value, child_node.token_extra_value]) child_node.node_value_len = len(child_node.token_mem_index_value) child_node.time_id = max(parent_node.time_id, child_node.time_id) @@ -576,8 +439,6 @@ def clear_tree_nodes(self): self.tree_total_tokens_num.arr[0] = 0 self.refed_tokens_num.arr[0] = 0 - self.swa_tree_total_pages_num = 0 - self.swa_refed_pages_num = 0 return def flush_cache(self): @@ -589,12 +450,6 @@ def dec_node_ref_counter(self, node: TreeNode): return # 如果减引用的是叶节点,需要先从 evict_tree_set 中移除 old_node = node - if ( - self._swa_pages_per_prompt_page > 0 - and old_node is not self.root_node - and self._node_direct_ref_num(old_node) == 1 - ): - self.swa_refed_pages_num -= self._node_last_valid_swa_pages_num(old_node) if old_node.is_leaf(): self.evict_tree_set.discard(old_node) @@ -614,12 +469,6 @@ def add_node_ref_counter(self, node: TreeNode): return # 如果减引用的是叶节点,需要先从 evict_tree_set 中移除 old_node = node - if ( - self._swa_pages_per_prompt_page > 0 - and old_node is not self.root_node - and self._node_direct_ref_num(old_node) == 0 - ): - self.swa_refed_pages_num += self._node_last_valid_swa_pages_num(old_node) if old_node.is_leaf(): self.evict_tree_set.discard(old_node) @@ -646,28 +495,12 @@ def get_mem_index_value_by_node(self, node: TreeNode) -> Optional[torch.Tensor]: ans_list.reverse() return torch.concat(ans_list, dim=0) - def get_extra_value_by_node(self, node: TreeNode): - if node is None or self.extra_value_ops is None: - return None - - ans_list = [] - while node is not None: - if node.token_extra_value is not None: - ans_list.append(node.token_extra_value) - node = node.parent - - ans_list.reverse() - return self._concat_extra(ans_list) - def get_refed_tokens_num(self): return self.refed_tokens_num.arr[0] def get_tree_total_tokens_num(self): return self.tree_total_tokens_num.arr[0] - def get_unrefed_swa_pages_num(self): - return self.swa_tree_total_pages_num - self.swa_refed_pages_num - def print_self(self, indent=0): self._print_helper(self.root_node, indent) @@ -682,88 +515,6 @@ def _print_helper(self, node: TreeNode, indent): self._print_helper(child, indent=indent + 2) return - def free_unreferenced_swa_pages(self, need_pages: int) -> None: - """DeepSeek-V4 swa free hook: 页 allocator 触底时回收未被命中终点保护的 swa 页。""" - if self.mem_manager is None or self.extra_value_ops is None: - return - allocator = self.mem_manager.swa_page_allocator - target = allocator.can_use_mem_size + int(need_pages) - planned_reclaim_pages = 0 - actual_reclaim_pages = 0 - while allocator.can_use_mem_size < target: - evict_slots = [] - invalidate_payloads = [] - evict_swa_pages = 0 - for free_last in (False, True): - visited = set() - for leaf in self.evict_tree_set: - if allocator.can_use_mem_size + evict_swa_pages >= target: - break - node = leaf - while node is not None and node is not self.root_node: - node_id = id(node) - if node_id in visited: - node = node.parent - continue - visited.add(node_id) - - has_direct_ref = self._node_direct_ref_num(node) > 0 - - payload = node.token_extra_value - if ( - len(node.token_mem_index_value) > 0 - and payload is not None - and payload.swa_page_valid is not None - ): - last_page = int(payload.swa_last_valid_page) - if last_page >= 0: - if free_last: - if has_direct_ref: - # 活跃请求恰好命中该节点时,最后一页仍需保留。 - node = node.parent - continue - page_slice = slice(last_page, last_page + 1) - else: - page_slice = slice(0, last_page) - valid_pages = int(payload.swa_page_valid[page_slice].sum().item()) - if valid_pages > 0: - start = page_slice.start * self.page_size - end = min(page_slice.stop * self.page_size, len(node.token_mem_index_value)) - if end > start: - evict_slots.append(node.token_mem_index_value[start:end]) - invalidate_payloads.append((payload, page_slice, free_last)) - evict_swa_pages += valid_pages * self._swa_pages_per_prompt_page - if allocator.can_use_mem_size + evict_swa_pages >= target: - break - node = node.parent - if allocator.can_use_mem_size + evict_swa_pages >= target: - break - if len(evict_slots) == 0: - break - free_pages_before = allocator.can_use_mem_size - self.mem_manager.evict_swa(torch.cat(evict_slots)) - for payload, page_slice, free_last in invalidate_payloads: - payload.swa_page_valid[page_slice] = False - if free_last: - payload.swa_last_valid_page = -1 - self.swa_tree_total_pages_num -= evict_swa_pages - planned_reclaim_pages += evict_swa_pages - actual_reclaim_pages += allocator.can_use_mem_size - free_pages_before - - # bitmap 是候选页账本,最终能否继续分配必须以物理 allocator 的真实水位为准。 - if allocator.can_use_mem_size < target or actual_reclaim_pages < planned_reclaim_pages: - logger.warning( - "DSV4 SWA reclaim mismatch: target_free_pages=%d actual_free_pages=%d " - "planned_reclaim_pages=%d actual_reclaim_pages=%d tree_pages=%d protected_pages=%d", - target, - allocator.can_use_mem_size, - planned_reclaim_pages, - actual_reclaim_pages, - self.swa_tree_total_pages_num, - self.swa_refed_pages_num, - ) - return - def free_radix_cache_to_get_enough_token(self, need_token_num): assert self.mem_manager is not None if need_token_num > self.mem_manager.allocator.can_use_mem_size: diff --git a/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py index 72b563699b..352c83fd7e 100644 --- a/lightllm/server/router/model_infer/infer_batch.py +++ b/lightllm/server/router/model_infer/infer_batch.py @@ -132,21 +132,16 @@ def add_reqs(self, requests: List[Tuple[int, int, Any, int]], init_prefix_cache: def free_a_req_mem(self, free_token_index: List, req: "InferReq"): if self.radix_cache is None: free_token_index.append(self.req_manager.req_to_token_indexs[req.req_idx][0 : req.hold_kv_len]) - if self.is_deepseek_v4: - # 槽位随 full 槽经 mem_manager.free 级联回收。pause 路径不释放 req_idx, - # 这里只复位出窗水位线;恢复从 0 重算,c128 request ring 会在读取前覆写。 - self.req_manager.init_compress_state(req.req_idx) elif CacheTier.GPU not in req.cache_tiers: self._free_req_mem_without_radix_insert(free_token_index=free_token_index, req=req) else: if not self.is_hybrid_att_model: - if self.is_deepseek_v4: - self._dsv4_full_att_free_req(free_token_index=free_token_index, req=req) - else: - self._full_att_free_req(free_token_index=free_token_index, req=req) + self._full_att_free_req(free_token_index=free_token_index, req=req) else: self._hybrid_att_free_req(free_token_index=free_token_index, req=req) assert len(req.hybrid_len_to_big_page_id) == 0 + if self.is_deepseek_v4: + self.req_manager.clear_runtime_state(req.req_idx) req.cur_kv_len = 0 req.hold_kv_len = 0 req.shm_req.shm_cur_kv_len = req.cur_kv_len @@ -192,40 +187,6 @@ def _full_att_free_req(self, free_token_index: List, req: "InferReq"): req.shared_kv_node = None return - def _dsv4_full_att_free_req(self, free_token_index: List, req: "InferReq"): - old_prefix_len = 0 if req.shared_kv_node is None else req.shared_kv_node.node_prefix_total_len - inserted_len = old_prefix_len - duplicate_prefix_len = old_prefix_len - - cache_len = self.radix_cache.align_len(req.cur_kv_len) - self.req_manager: DeepseekV4ReqManager - if cache_len > old_prefix_len: - payload = self.req_manager.build_prompt_cache_payload(cache_len) - value = self.req_manager.req_to_token_indexs[req.req_idx][:cache_len].detach().cpu() - - payload.swa_page_valid = self.req_manager.swa_page_valid_from_watermark(req.req_idx, cache_len) - payload.refresh_swa_last_valid_page() - - key = torch.tensor(req.get_input_token_ids()[0:cache_len], dtype=torch.int64, device="cpu") - duplicate_prefix_len, _ = self.radix_cache.insert(key, value[:cache_len], extra_value=payload) - inserted_len = cache_len - - dense_row = self.req_manager.req_to_token_indexs[req.req_idx] - if duplicate_prefix_len > old_prefix_len: - free_token_index.append(dense_row[old_prefix_len:duplicate_prefix_len]) - if req.hold_kv_len > inserted_len: - free_token_index.append(dense_row[inserted_len : req.hold_kv_len]) - if len(free_token_index) == 0: - free_token_index.append(dense_row[0:0]) - - self.req_manager.init_compress_state(req.req_idx) - - if req.shared_kv_node is not None: - assert req.shared_kv_node.node_prefix_total_len <= max(inserted_len, old_prefix_len) - self.radix_cache.dec_node_ref_counter(req.shared_kv_node) - req.shared_kv_node = None - return - def _hybrid_att_free_req(self, free_token_index: List, req: "InferReq"): assert g_infer_context.is_hybrid_att_model is True args = get_env_start_args() @@ -452,44 +413,47 @@ def get_can_alloc_token_num(self): return self.req_manager.mem_manager.allocator.can_use_mem_size + radix_cache_unref_token_num def get_can_alloc_dsv4_swa_page_num(self): - pages = int(self.req_manager.mem_manager.swa_page_allocator.can_use_mem_size) - if self.radix_cache is not None: - pages += self.radix_cache.get_unrefed_swa_pages_num() - return pages + return int(self.req_manager.mem_manager.swa_page_allocator.can_use_mem_size) def save_hybrid_state_to_cache(self, b_req_idx: torch.Tensor, reqs: List["InferReq"]): """Snapshot request-level attention state at big/small-page boundaries.""" - if not self.is_hybrid_att_model: + if not self.is_hybrid_att_model or self.radix_cache is None: return # Request-state snapshot at a big-page boundary. big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num big_page_buffer_ids = [] - for req in reqs: - cur_input_len = req.get_chuncked_input_token_len() - if cur_input_len % big_page_token_num == 0 and cur_input_len <= req.hybrid_cache_len: + checkpoint_lens = [] + req_rows = [] + for row, req in enumerate(reqs): + chunk_end = req.get_chuncked_input_token_len() + first_boundary = (req.cur_kv_len // big_page_token_num + 1) * big_page_token_num + for length in range(first_boundary, min(chunk_end, req.hybrid_cache_len) + 1, big_page_token_num): big_page_id = self.radix_cache.big_page_buffers.alloc_one_state_cache() assert big_page_id is not None + assert length not in req.hybrid_len_to_big_page_id + req.hybrid_len_to_big_page_id[length] = big_page_id big_page_buffer_ids.append(big_page_id) - assert cur_input_len not in req.hybrid_len_to_big_page_id - req.hybrid_len_to_big_page_id[cur_input_len] = big_page_id - else: - big_page_buffer_ids.append(-1) + checkpoint_lens.append(length) + req_rows.append(row) - assert len(b_req_idx) == len(big_page_buffer_ids) - if any(buffer_id != -1 for buffer_id in big_page_buffer_ids): + if big_page_buffer_ids: + selected_rows = torch.tensor(req_rows, dtype=torch.int64, device=b_req_idx.device) self.req_manager.save_big_page_states( - b_req_idx=b_req_idx, - req_indexes=[req.req_idx for req in reqs], + b_req_idx=b_req_idx[selected_rows], + req_indexes=[reqs[row].req_idx for row in req_rows], buffer_indexes=big_page_buffer_ids, + checkpoint_lens=checkpoint_lens, ) - assert not self.args.disable_chunked_prefill, "chunked prefill must be enabled for hybrid attention models" + assert ( + self.is_deepseek_v4 or not self.args.disable_chunked_prefill + ), "chunked prefill must be enabled for linear attention models" # Request-state snapshot at the final small-page boundary. for req in reqs: # 判断本次prefill 完以后 kv 的长度是否到达 hybrid checkpoint 的存储边界。 - if req.get_chuncked_input_token_len() == req.hybrid_cache_len: + if req.cur_kv_len < req.hybrid_cache_len <= req.get_chuncked_input_token_len(): assert req.tail_small_page_buffer_id is None if req.hybrid_cache_len % big_page_token_num != 0: self.radix_cache.free_one_small_page_buffer() @@ -500,6 +464,7 @@ def save_hybrid_state_to_cache(self, b_req_idx: torch.Tensor, reqs: List["InferR req_idx=req.req_idx, buffer_idx=dst_buffer_idx, state_cache_manager=self.radix_cache.small_page_buffers, + checkpoint_len=req.hybrid_cache_len, ) return @@ -684,8 +649,6 @@ def _init_all_state(self): self.final_token_metadata = FinalTokenMetadataExt(self) g_infer_context.req_manager.req_sampling_params_manager.init_req_sampling_params(self) - if hasattr(g_infer_context.req_manager, "init_compress_state"): - g_infer_context.req_manager.init_compress_state(req_idx=self.req_idx) self.stop_sequences = self.sampling_param.shm_param.stop_sequences.to_list() self.multimodal_params = self.multimodal_params.to_dict() @@ -702,6 +665,11 @@ def _init_all_state(self): if g_infer_context.is_hybrid_att_model: block_num = self.shm_req.hybrid_token_hash_list.size self.hybrid_cache_len = block_num * self.args.linear_att_hash_page_size + for image_start, image_end in reversed(self.image_block_spans): + if image_start < self.hybrid_cache_len < image_end: + self.hybrid_cache_len = ( + image_start // self.args.linear_att_hash_page_size * self.args.linear_att_hash_page_size + ) self.hybrid_len_to_big_page_id = SortedDict() return @@ -715,25 +683,12 @@ def _match_radix_cache(self): input_token_ids = self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] key = torch.tensor(input_token_ids, dtype=torch.int64, device="cpu") key = key[0 : len(key) - 1] # 最后一个不需要,因为需要一个额外的token,让其在prefill的时候输出下一个token的值 - # DSV4 prompt cache may reclaim earlier SWA pages, so an image-internal hit must be recomputed. - if g_infer_context.is_deepseek_v4 and self.image_block_spans: - while True: - _, matched_len, _ = g_infer_context.radix_cache.match_prefix(key, update_refs=False) - for image_start, image_end in self.image_block_spans: - if image_start < matched_len < image_end: - key = key[:image_start] - break - else: - break share_node, kv_len, value_tensor = g_infer_context.radix_cache.match_prefix(key, update_refs=True) if share_node is not None: self.shared_kv_node = share_node ready_cache_len = share_node.node_prefix_total_len # 从 cpu 到 gpu 是流内阻塞操作 g_infer_context.req_manager.req_to_token_indexs[self.req_idx, 0:ready_cache_len] = value_tensor - # DeepSeek-V4 命中无需恢复 compressor 状态: 槽位由 full_to_* 映射键控(radix - # 持有 full 槽即有效),c4 状态随 swa 页常驻;命中点按 256 token 对齐,c128 - # request ring 的下一组会在首次读取前完整覆写。 self.cur_kv_len = int(ready_cache_len) # 序列化问题, 该对象可能为numpy.int64,用 int(*)转换 self.hold_kv_len = self.cur_kv_len assert self.hold_kv_len % self.args.page_size == 0 @@ -747,6 +702,8 @@ def _hybrid_match_radix_cache(self): g_infer_context.is_hybrid_att_model is True ), "current _hybrid_match_radix_cache only supports hybrid attention models, to do..." enable_prompt_cache = (not self.sampling_param.disable_prompt_cache) and g_infer_context.radix_cache is not None + if g_infer_context.is_deepseek_v4: + enable_prompt_cache = enable_prompt_cache and g_infer_context.get_can_alloc_dsv4_swa_page_num() >= 2 block_hashs = self.shm_req.hybrid_token_hash_list.get_all() hash_page_size = self.args.linear_att_hash_page_size match_tokens = min(len(block_hashs) * hash_page_size, self.get_cur_total_len() - 1) @@ -761,6 +718,18 @@ def _hybrid_match_radix_cache(self): input_token_ids = self.shm_req.shm_prompt_ids.arr[0 : self.get_cur_total_len()] key = torch.tensor(input_token_ids[0:match_tokens], dtype=torch.int64, device="cpu") assert len(key) == len(block_hashs) * hash_page_size + if self.image_block_spans: + while True: + _, matched_len, _ = g_infer_context.radix_cache.match_prefix( + key, block_hashs=block_hashs, update_refs=False + ) + for image_start, image_end in self.image_block_spans: + if image_start < matched_len < image_end: + limit = image_start // hash_page_size * hash_page_size + key, block_hashs = key[:limit], block_hashs[: limit // hash_page_size] + break + else: + break share_node, kv_len, value_tensor = g_infer_context.radix_cache.match_prefix( key, block_hashs=block_hashs, update_refs=True ) @@ -777,7 +746,7 @@ def _hybrid_match_radix_cache(self): assert self.tail_small_page_buffer_id is None # 恢复 hybrid checkpoint g_infer_context.req_manager.restore_big_page_state( - big_page_buffer_idx=share_node.big_page_buffer_idx, req=self + big_page_buffer_idx=share_node.big_page_buffer_idx, req=self, checkpoint_len=ready_cache_len ) else: # 小页匹配 @@ -793,6 +762,7 @@ def _hybrid_match_radix_cache(self): # 恢复 hybrid checkpoint g_infer_context.req_manager.restore_small_page_state( req=self, + checkpoint_len=ready_cache_len, ) else: # 如果 大页本质是被启用的,则需要使用小页的匹配结果, 将小页的kv 复制到的新申请的kv位置,同时释放 @@ -827,6 +797,7 @@ def _hybrid_match_radix_cache(self): self.shared_kv_node = share_node # 只是为了保证 restore_small_page_state 正确调用 g_infer_context.req_manager.restore_small_page_state( req=self, + checkpoint_len=shared_kv_len, ) self.shared_kv_node = None @@ -851,7 +822,9 @@ def _hybrid_match_radix_cache(self): assert self.tail_small_page_buffer_id is None # 恢复 hybrid checkpoint g_infer_context.req_manager.restore_big_page_state( - big_page_buffer_idx=share_node.big_page_buffer_idx, req=self + big_page_buffer_idx=share_node.big_page_buffer_idx, + req=self, + checkpoint_len=ready_cache_len, ) self.shm_req.shm_cur_kv_len = self.cur_kv_len @@ -923,32 +896,26 @@ def get_chuncked_input_token_ids(self): return self.shm_req.shm_prompt_ids.arr[0 : self._get_chunked_input_end()] def get_chuncked_input_token_ids_for_hybrid_att(self): - big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num - - chunked_start = self.cur_kv_len - chunked_end = chunked_start + self.args.chunked_prefill_size - big_page_end = ((chunked_start // big_page_token_num) + 1) * big_page_token_num - total_end = self.get_cur_total_len() - end = min(total_end, chunked_end, big_page_end) - - if chunked_start < self.hybrid_cache_len < end: - # hybrid checkpoint 对应需要存储的部分。 - end = self.hybrid_cache_len - - return self.shm_req.shm_prompt_ids.arr[0:end] + return self.shm_req.shm_prompt_ids.arr[0 : self.get_chuncked_input_token_len_for_hybrid_att()] def get_chuncked_input_token_len(self): return self._get_chunked_input_end() def get_chuncked_input_token_len_for_hybrid_att(self): + end = self._get_chunked_input_end() + if g_infer_context.is_deepseek_v4 and self.args.disable_chunked_prefill: + return end + if g_infer_context.radix_cache is None: + return end big_page_token_num = self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num - chunked_start = self.cur_kv_len - chunked_end = chunked_start + self.args.chunked_prefill_size - big_page_end = ((chunked_start // big_page_token_num) + 1) * big_page_token_num - total_end = self.get_cur_total_len() - end = min(total_end, chunked_end, big_page_end) - if chunked_start < self.hybrid_cache_len < end: + big_page_end = (self.cur_kv_len // big_page_token_num + 1) * big_page_token_num + end = min(end, big_page_end) + if self.cur_kv_len < self.hybrid_cache_len < end: end = self.hybrid_cache_len + for image_start, image_end in self.image_block_spans: + if image_start < end < image_end: + end = image_start if self.cur_kv_len < image_start else image_end + break return end def set_next_gen_token_id(self, next_token_id: int, logprob: float, output_len: int, rank: int = -1): @@ -1060,11 +1027,7 @@ def get_dsv4_prefill_need_swa_page_num(self, is_chuncked_prefill: bool) -> int: if end <= start: return 0 - first_new_page = (start + self.dsv4_swa_page_size - 1) // self.dsv4_swa_page_size - last_page = (end - 1) // self.dsv4_swa_page_size - swa_page_num = last_page - first_new_page + 1 - - return swa_page_num + return g_infer_context.req_manager.get_swa_page_need(self.req_idx, start, end) def get_dsv4_recover_need_swa_page_num(self) -> int: swa_page_num = self.get_dsv4_prefill_need_swa_page_num(is_chuncked_prefill=False) @@ -1090,23 +1053,9 @@ def get_dsv4_decode_need_swa_page_num(self) -> int: if seq_len <= 0: return 0 - swa_page_num = 0 - # Main model prepares current token plus draft-verify rows: SWA + compressed slots. - for step in range(self.mtp_step + 1): - cur_seq_len = seq_len + step - if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: - swa_page_num += 1 - if self.args.mtp_mode == "dspark": - # DSpark proposal allocates one private SWA scratch page per request. - swa_page_num += 1 - else: - # EAGLE draft forwards after the first one consume newly appended draft-only rows. - # The DeepSeek-V4 MTP draft layer is compress_ratio=0, so these rows need only SWA. - for step in range(self.mtp_step + 1, self.mtp_step * 2): - cur_seq_len = seq_len + step - if (cur_seq_len - 1) % self.dsv4_swa_page_size == 0: - swa_page_num += 1 - return swa_page_num + width = self.mtp_step + 1 if self.args.mtp_mode == "dspark" else max(self.mtp_step + 1, self.mtp_step * 2) + need = g_infer_context.req_manager.get_swa_page_need(self.req_idx, seq_len - 1, seq_len - 1 + width) + return need + int(self.args.mtp_mode == "dspark") class InferReqUpdatePack: diff --git a/lightllm/server/router/model_infer/mode_backend/base_backend.py b/lightllm/server/router/model_infer/mode_backend/base_backend.py index edf16809ea..0bca04e4c7 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -16,7 +16,7 @@ from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( disable_eplb_model_init, ) -from lightllm.common.req_manager import DeepseekV4ReqManager, ReqManagerForMamba +from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager from lightllm.common.req_manager import HybridAttentionReqManager @@ -174,26 +174,12 @@ def init_model(self, kvargs): small_page_buffers=self.small_page_buffers, ) else: - radix_page_size = self.args.page_size - radix_extra_value_ops = None - if self.is_deepseek_v4: - radix_page_size = self.model.req_manager.get_prompt_cache_page_size() - radix_extra_value_ops = self.model.req_manager.get_prompt_cache_value_ops() self.radix_cache = RadixCache( total_token_num=self.model.mem_manager.size, rank_in_node=self.rank_in_node, mem_manager=self.model.mem_manager, - page_size=radix_page_size, - extra_value_ops=radix_extra_value_ops, + page_size=self.args.page_size, ) - if self.is_deepseek_v4: - self.model.mem_manager.register_swa_free_hook(self.radix_cache.free_unreferenced_swa_pages) - - if not self.disable_chunked_prefill and radix_page_size > 1: - assert self.args.chunked_prefill_size % radix_page_size == 0, ( - f"chunked_prefill_size={self.args.chunked_prefill_size} must be divisible by " - f"prompt-cache page_size={radix_page_size}" - ) if "prompt_cache_kv_buffer" in model_cfg: assert self.use_dynamic_prompt_cache diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index f865a637a2..ef9cf5c3f2 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -51,7 +51,7 @@ def init_dsv4_cache_transfer(self, mem_managers: List[MemoryManager]) -> None: mem_manager.c4_pool.buffer.data_ptr() if has_c4 else 0, mem_manager.c4_indexer_pool.buffer.data_ptr() if has_c4 else 0, mem_manager.c128_pool.buffer.data_ptr() if has_c128 else 0, - mem_manager.full_to_swa_indexs.data_ptr(), + mem_manager.req_to_swa_pages.data_ptr(), mem_manager.swa_pool.buffer.data_ptr(), mem_manager.c4_state_buffer.data_ptr() if has_c4 else 0, mem_manager.c4_indexer_state_buffer.data_ptr() if has_c4 else 0, @@ -139,11 +139,64 @@ def build_shared_kv_trans_tasks( max_kv_len_dp_rank=int(max_kv_len_dp_rank), max_kv_len_mem_manager_index=int(max_kv_len_mem_manager_index), max_kv_len_mem_indexes=max_kv_len_mem_indexes, + source_req_idx=max_kv_len_req_idx, ) ) return trans_tasks + def _transfer_dsv4_checkpoints(self, trans_tasks): + """Transfer CPU snapshots for the history fetched from another DP rank. + + CUDA IPC exports the packed history and private runtime only. NCCL moves + the checkpoint bytes through temporary GPU tensors so pinned host pools + keep their process-local ownership and registration. + """ + args = self.backend.args + big_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + if big_tokens > args.max_req_total_len: + return + group = self.backend.node_nccl_group + rank = dist.get_rank(group=group) + local_tasks = [ + ( + task.max_kv_len_mem_manager_index, + task.req.req_id, + task.req.cur_kv_len, + task.req.cur_kv_len + len(task.mem_indexes), + ) + for task in trans_tasks + ] + all_tasks = [None for _ in range(dist.get_world_size(group=group))] + dist.all_gather_object(all_tasks, local_tasks, group=group) + buffers = self.backend.model.mem_manager.big_page_buffers + for destination, tasks in enumerate(all_tasks): + for source, req_id, start, end in tasks: + first = (start // big_tokens + 1) * big_tokens + lengths = list(range(first, end + 1, big_tokens)) + if not lengths or rank not in (source, destination): + continue + req = g_infer_context.requests_mapping[req_id] + if rank == source: + shared_ids = self.backend.radix_cache.get_big_page_ids_by_node(req.shared_kv_node) + indexes = [ + req.hybrid_len_to_big_page_id[length] + if length in req.hybrid_len_to_big_page_id + else shared_ids[length // big_tokens - 1] + for length in lengths + ] + for index in indexes: + staging = buffers.buffer[index].cuda(non_blocking=True) + dist.send(staging, dst=dist.get_global_rank(group, destination), group=group) + else: + staging = torch.empty((buffers.buffer.shape[1],), dtype=torch.uint8, device="cuda") + for length in lengths: + index = buffers.alloc_one_state_cache() + assert index is not None + dist.recv(staging, src=dist.get_global_rank(group, source), group=group) + buffers.buffer[index].copy_(staging, non_blocking=True) + req.hybrid_len_to_big_page_id[length] = index + def kv_trans(self, trans_tasks: List["TransTask"]): # kv 传输 if len(trans_tasks) > 0: @@ -151,9 +204,7 @@ def kv_trans(self, trans_tasks: List["TransTask"]): req_manager = g_infer_context.req_manager prompt_cache_page_size = req_manager.get_prompt_cache_page_size() req_list = [] - ready_list = [] seq_list = [] - dst_full_slot_views = [] task_meta_data = [] history_block_nums = [] for trans_task in trans_tasks: @@ -162,29 +213,21 @@ def kv_trans(self, trans_tasks: List["TransTask"]): dst_full_slots = req_manager.req_to_token_indexs[trans_task.req.req_idx, start:end] dst_full_slots.copy_(trans_task.mem_indexes, non_blocking=True) req_list.append(trans_task.req.req_idx) - ready_list.append(start) seq_list.append(end) - dst_full_slot_views.append(dst_full_slots) task_meta_data.extend( [ trans_task.max_kv_len_mem_manager_index, - end - start, trans_task.max_kv_len_mem_indexes.data_ptr(), dst_full_slots.data_ptr(), + trans_task.source_req_idx, + trans_task.req.req_idx, + end, ] ) history_block_nums.append((end - start) // prompt_cache_page_size) - # Keep full slots in the same request-major order as req_list. - new_full_slots = ( - dst_full_slot_views[0] if len(dst_full_slot_views) == 1 else torch.cat(dst_full_slot_views) - ) - req_manager.prepare_pd_decode_cache( - req_list=req_list, - ready_list=ready_list, - seq_list=seq_list, - new_full_slots=new_full_slots, - ) + for req_idx, end in zip(req_list, seq_list): + req_manager.prepare_swa(req_idx, end - prompt_cache_page_size, end) # The history kernel consumes (task index, block index) pairs in block-major order. history_meta_data = [] @@ -194,7 +237,8 @@ def kv_trans(self, trans_tasks: List["TransTask"]): history_meta_data.extend([task_index, block_index]) # transfer_meta packs two flat uint64 tables into one H2D copy: - # task_meta: (source manager, token count, source slots pointer, destination slots pointer) + # task_meta: (source manager, source/destination slots pointers, + # source/destination request IDs, logical end) # history_meta: (task index, block index) # For two tasks with 2 and 1 history blocks, the layout is: # [task0 fields, task1 fields, (0, 0), (1, 0), (0, 1)] @@ -237,6 +281,24 @@ def kv_trans(self, trans_tasks: List["TransTask"]): transfer_token_num = sum(len(trans_task.mem_indexes) for trans_task in trans_tasks) self.backend.logger.info(f"dp_i {self.dp_rank_in_node} transfer kv tokens num: {transfer_token_num}") + if self.backend.is_deepseek_v4: + self._transfer_dsv4_checkpoints(trans_tasks) + args = self.backend.args + big_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + for task in trans_tasks: + req = task.req + end = req.cur_kv_len + len(task.mem_indexes) + if end == req.hybrid_cache_len and end % big_tokens and self.backend.radix_cache is not None: + self.backend.radix_cache.free_one_small_page_buffer() + req.tail_small_page_buffer_id = self.backend.small_page_buffers.alloc_one_state_cache() + if req.tail_small_page_buffer_id is not None: + g_infer_context.req_manager.save_state( + req.req_idx, + req.tail_small_page_buffer_id, + self.backend.small_page_buffers, + checkpoint_len=end, + ) + if self.backend.is_deepseek_v4 and self.backend.args.enable_cpu_cache: # CPU-cache restore can evict source radix pages before the scheduler all-gather fences this stream. dist.barrier(group=self.backend.node_nccl_group) @@ -255,3 +317,4 @@ class TransTask: max_kv_len_dp_rank: int max_kv_len_mem_manager_index: int max_kv_len_mem_indexes: torch.Tensor + source_req_idx: int diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py index 5c1746474b..34f998230a 100644 --- a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py @@ -150,7 +150,10 @@ def _submit_dsv4_store_batch( source_mem_indexes = slot.source_mem_indexes[:page_num] torch.stack([item.source_mem_indexes for item in store_pages], out=source_mem_indexes) staging = slot.buffer[:page_num] - operator.pack_cpu_cache_pages(source_mem_indexes, staging) + source_req_meta = torch.tensor( + [[item.req_idx, item.checkpoint_len] for item in store_pages], dtype=torch.int32, device="cuda" + ) + operator.pack_cpu_cache_pages(source_mem_indexes, source_req_meta, staging) pack_event = torch.cuda.Event() pack_event.record() @@ -227,6 +230,8 @@ def store_completed_prefill_pages( session=session, cpu_page_index=cpu_page_index, source_mem_indexes=source_mem_indexes, + req_idx=req.req_idx, + checkpoint_len=token_start + token_page_size, ) ) if session.disabled or session.next_page_index >= len(token_hashes): @@ -292,13 +297,9 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): if loadable_end != 0: token_num = loadable_end - gpu_kv_len full_need = token_num - swa_need = 2 if self.backend.radix_cache is not None: radix_cache = self.backend.radix_cache radix_cache.free_radix_cache_to_get_enough_token(full_need) - swa_shortage = swa_need - int(mem_manager.swa_page_allocator.can_use_mem_size) - if swa_shortage > 0: - radix_cache.free_unreferenced_swa_pages(swa_shortage) loadable_end = mem_manager.get_loadable_cpu_cache_end( gpu_kv_len, @@ -315,20 +316,55 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): first_page_index = gpu_kv_len // layout.token_page_size cpu_pages = page_list[first_page_index : loaded_end // layout.token_page_size] page_indexes_cuda = torch.tensor(cpu_pages, dtype=torch.int32, device="cuda") - plan = mem_manager.prepare_cpu_cache_load(token_num=token_num, loaded_end=loaded_end) - mem_manager.operator.load_cpu_cache_pages( - plan=plan, - page_indexes=page_indexes_cuda, - cpu_cache_client=self.cpu_cache_client, - first_page_history_offset_tokens=gpu_kv_len % layout.token_page_size, + req_manager = self.backend.model.req_manager + req_manager.prepare_swa(req.req_idx, loaded_end - 256, loaded_end) + resume_slots = req_manager.get_swa_slots( + req.req_idx, torch.arange(loaded_end - 256, loaded_end, device="cuda") + ) + plan = mem_manager.prepare_cpu_cache_load( + token_num=token_num, loaded_end=loaded_end, resume_swa_slots=resume_slots ) - mem_manager.commit_cpu_cache_load_plan(plan) + try: + mem_manager.operator.load_cpu_cache_pages( + plan=plan, + page_indexes=page_indexes_cuda, + cpu_cache_client=self.cpu_cache_client, + first_page_history_offset_tokens=gpu_kv_len % layout.token_page_size, + ) + except Exception: + mem_manager.free(plan.mem_indexes) + req_manager.clear_runtime_state(req.req_idx) + raise self.backend.model.req_manager.req_to_token_indexs[ req.req_idx, gpu_kv_len:loaded_end ] = plan.mem_indexes - self.backend.model.req_manager.finish_cpu_cache_load(req.req_idx, loaded_end) req.cur_kv_len = loaded_end req.hold_kv_len = loaded_end + if self.backend.radix_cache is not None: + big_page_tokens = ( + self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num + ) + first_boundary = (gpu_kv_len // big_page_tokens + 1) * big_page_tokens + for boundary in range(first_boundary, loaded_end + 1, big_page_tokens): + buffer_idx = mem_manager.big_page_buffers.alloc_one_state_cache() + assert buffer_idx is not None + page_idx = page_list[boundary // layout.token_page_size - 1] + mem_manager.big_page_buffers.buffer[buffer_idx].copy_( + self.cpu_cache_client.cpu_kv_cache_tensor[page_idx, layout.swa_offset :] + ) + req.hybrid_len_to_big_page_id[boundary] = buffer_idx + if loaded_end == req.hybrid_cache_len and loaded_end % big_page_tokens: + self.backend.radix_cache.free_one_small_page_buffer() + req.tail_small_page_buffer_id = ( + req_manager.small_page_buffers.alloc_one_state_cache() + ) + if req.tail_small_page_buffer_id is not None: + req_manager.save_state( + req.req_idx, + req.tail_small_page_buffer_id, + req_manager.small_page_buffers, + checkpoint_len=loaded_end, + ) idle_token_num -= token_num if is_master_in_dp: @@ -400,6 +436,8 @@ class Dsv4StorePage: cpu_page_index: int # 该 checkpoint page 对应的 GPU KV slot 编号。 source_mem_indexes: torch.Tensor + req_idx: int + checkpoint_len: int @dataclasses.dataclass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index f404cf66b2..f8f0045fb9 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -70,7 +70,10 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: # pending 期间优先重新匹配 radix;准入失败时释放引用,留待下轮重试。 if req_obj.pd_task_num == 0 and not req_obj.infer_aborted: - req_obj._match_radix_cache() + if g_infer_context.is_hybrid_att_model: + req_obj._hybrid_match_radix_cache() + else: + req_obj._match_radix_cache() if not self._decode_node_gen_trans_tasks(req_obj=req_obj): PDDecodeNode._drop_pending_prompt_cache(self, req_obj) continue @@ -91,6 +94,8 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: continue if req_obj.pd_task_failed_num > 0: + if isinstance(self.model.req_manager, DeepseekV4ReqManager): + self.model.req_manager.clear_runtime_state(req_obj.req_idx) # KV 传输失败:强制补 finish token 并结束。 # abort 优先标 ABORTED;纯传输错误标 ERROR(不再误用 STOP)。 if not req_obj.finish_status.is_finished(): @@ -127,12 +132,41 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: if self.is_master_in_dp: req_obj.shm_req.shm_cur_kv_len = req_obj.cur_kv_len + if ( + isinstance(self.model.req_manager, DeepseekV4ReqManager) + and req_obj.pd_task_failed_num == 0 + and not req_obj.infer_aborted + and req_obj.cur_kv_len == req_obj.shm_req.input_len + and req_obj.hybrid_cache_len > 0 + and req_obj.hybrid_cache_len == (req_obj.shm_req.input_len - 1) // 256 * 256 + and req_obj.tail_small_page_buffer_id is None + and self.radix_cache is not None + ): + # The packed transfer includes both the arbitrary live tail and + # the last aligned continuation. Capture it before decode advances. + self.radix_cache.free_one_small_page_buffer() + req_obj.tail_small_page_buffer_id = self.small_page_buffers.alloc_one_state_cache() + if req_obj.tail_small_page_buffer_id is not None: + self.model.req_manager.save_state( + req_obj.req_idx, + req_obj.tail_small_page_buffer_id, + self.small_page_buffers, + checkpoint_len=req_obj.hybrid_cache_len, + ) + ans_list.append(req_obj) return ans_list def _drop_pending_prompt_cache(self, req_obj: InferReq) -> None: """准入失败时撤销 D 侧命中,避免 pending 请求长期占用 radix 引用。""" assert req_obj.pd_task_num == 0 + shared_len = 0 if req_obj.shared_kv_node is None else req_obj.shared_kv_node.node_prefix_total_len + if req_obj.hold_kv_len > shared_len: + self.model.mem_manager.free( + self.model.req_manager.req_to_token_indexs[req_obj.req_idx, shared_len : req_obj.hold_kv_len] + ) + if isinstance(self.model.req_manager, DeepseekV4ReqManager): + self.model.req_manager.clear_runtime_state(req_obj.req_idx) if req_obj.shared_kv_node is not None: self.radix_cache.dec_node_ref_counter(req_obj.shared_kv_node) req_obj.shared_kv_node = None @@ -162,7 +196,6 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: is_dsv4_req_manager = isinstance(req_manager, DeepseekV4ReqManager) if is_dsv4_req_manager: mem_manager = req_manager.mem_manager - ready_len = req_obj.cur_kv_len if need_mem_size > g_infer_context.get_can_alloc_token_num(): return False @@ -172,14 +205,9 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: return False prompt_page = req_manager.get_prompt_cache_page_size() - swa_start = max(ready_len, max(0, input_len // prompt_page * prompt_page - prompt_page)) - swa_page = mem_manager.swa_pool.page_size - swa_need = max(0, (input_len - 1) // swa_page - (swa_start + swa_page - 1) // swa_page + 1) + swa_start = max(0, (input_len - 1) // prompt_page * prompt_page - prompt_page) + swa_need = req_manager.get_swa_page_need(req_obj.req_idx, swa_start, input_len) - if self.radix_cache is not None: - swa_shortage = swa_need - mem_manager.swa_page_allocator.can_use_mem_size - if swa_shortage > 0: - self.radix_cache.free_unreferenced_swa_pages(swa_shortage) if swa_need > mem_manager.swa_page_allocator.can_use_mem_size: return False @@ -192,9 +220,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: if is_dsv4_req_manager: req_manager.prepare_pd_decode_cache( req_list=[req_obj.req_idx], - ready_list=[req_obj.cur_kv_len], seq_list=[input_len], - new_full_slots=mem_indexes, ) torch.cuda.current_stream().synchronize() @@ -219,7 +245,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: # 混合注意力模型还需接收请求运行态 buffer(如 linear attention 的 conv/SSM 状态)。 # 通过本地 req_idx 定位运行态 buffer 的恢复位置。 - if g_infer_context.is_hybrid_att_model: + if g_infer_context.is_hybrid_att_model and not is_dsv4_req_manager: self._create_pd_trans_task( req_obj=req_obj, mem_indexes=[], diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py index 336266c531..82db081e91 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_impl.py @@ -88,7 +88,7 @@ def _prefill_chuncked_handle_func( break if prefill_finished and len(trans_task_list) != 0 and output_len == 1: - if g_infer_context.is_hybrid_att_model: + if g_infer_context.is_hybrid_att_model and not self.is_deepseek_v4: # 混合注意力模型除 KV 外,还需传输 prefill 完成时的请求运行态 buffer(如 linear attention 的 conv/SSM 状态)。 trans_task_list.append( self._create_pd_trans_task( diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index 30b863156d..d2632c42cf 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -549,7 +549,7 @@ def is_linear_att_mixed_model(model_path: str) -> bool: def is_hybrid_att_model(model_path: str) -> bool: """Models whose non-full attention state follows hybrid checkpoint pages.""" - return is_linear_att_mixed_model(model_path) + return get_model_type(model_path) == "deepseek_v4" or is_linear_att_mixed_model(model_path) def get_model_type(model_path: str) -> Optional[str]: diff --git a/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py index b6041c6e38..b4a28ee48f 100644 --- a/lightllm/utils/kv_cache_utils.py +++ b/lightllm/utils/kv_cache_utils.py @@ -17,12 +17,11 @@ ) from lightllm.utils.log_utils import init_logger from lightllm.utils.config_utils import ( - get_deepseek_v4_compress_rates, - get_config_json, get_num_key_value_heads, get_head_dim, get_layer_num, is_hybrid_att_model, + get_model_type, ) from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class from lightllm.common.kv_cache_mem_manager import ( @@ -70,7 +69,7 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": args = get_env_start_args() assert args.enable_cpu_cache - is_hybrid_model = is_hybrid_att_model(args.model_dir) + is_hybrid_model = is_hybrid_att_model(args.model_dir) and get_model_type(args.model_dir) != "deepseek_v4" mem_manager_class = None if is_hybrid_model else select_mem_manager_class() if is_hybrid_model: hybrid_config = get_hybrid_cache_config() @@ -118,16 +117,8 @@ def calcu_cpu_cache_meta() -> "CpuKVCacheMeta": scale_data_type=get_llm_data_type(), ) elif mem_manager_class is DeepseekV4MemoryManager: - from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4CpuCacheLayout - - config = get_config_json(args.model_dir) - layer_num = get_layer_num(args.model_dir) + get_added_mtp_kv_layer_num() - layout = DeepseekV4CpuCacheLayout.from_compress_rates( - compress_rates=get_deepseek_v4_compress_rates(config, layer_num), - token_page_size=args.cpu_cache_token_page_size, - head_dim=get_head_dim(args.model_dir), - indexer_head_dim=config["index_head_dim"], - ) + layout = get_hybrid_cache_config() + layer_num = layout.layer_num cpu_cache_meta = CpuKVCacheMeta( page_num=0, token_page_size=args.cpu_cache_token_page_size, diff --git a/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py index 206438d986..dbb8e02cdd 100644 --- a/test/unit/test_deepseek_v4_dspark.py +++ b/test/unit/test_deepseek_v4_dspark.py @@ -8,10 +8,12 @@ from lightllm.utils.config_utils import get_deepseek_v4_compress_rates -def test_dspark_decode_admission_reserves_one_swa_scratch_page(): - from lightllm.server.router.model_infer.infer_batch import InferReq +def test_dspark_decode_admission_reserves_one_swa_scratch_page(monkeypatch): + from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context req = InferReq.__new__(InferReq) + req.req_idx = 0 + monkeypatch.setattr(g_infer_context, "req_manager", SimpleNamespace(get_swa_page_need=lambda req, start, end: 0)) req.args = SimpleNamespace(mtp_mode="eagle") req.mtp_step = 1 req.dsv4_swa_page_size = 128 @@ -193,16 +195,7 @@ def test_build_dspark_swa_index_exposes_history_and_complete_block(): block_size = 3 window = 4 - req_to_token = torch.tensor( - [ - list(range(0, 10)), - list(range(10, 20)), - [20] * 10, - ], - dtype=torch.int32, - device="cuda", - ) - full_to_swa = torch.arange(21, dtype=torch.int32, device="cuda") + 100 + req_to_swa_pages = torch.tensor([[1], [2], [3]], dtype=torch.int32, device="cuda") req_idx = torch.tensor([0, 0, 0, 1, 1, 1], dtype=torch.int32, device="cuda") positions = torch.tensor([4, 5, 6, 2, 3, 4], dtype=torch.int32, device="cuda") padded_width = 8 @@ -214,8 +207,7 @@ def test_build_dspark_swa_index_exposes_history_and_complete_block(): build_dspark_swa_index( req_idx=req_idx, positions=positions, - req_to_token_indexs=req_to_token, - full_to_swa_indexs=full_to_swa, + req_to_swa_pages=req_to_swa_pages, scratch_pages=scratch_pages, swa_index=indices, swa_length=lengths, @@ -224,18 +216,17 @@ def test_build_dspark_swa_index_exposes_history_and_complete_block(): block_size=block_size, page_size=128, hold_req_id=2, - hold_full_slot=20, hold_swa_slot=120, ) expected_indices = torch.tensor( [ - [103, 102, 101, 100, 256, 257, 258, -1], - [103, 102, 101, 100, 256, 257, 258, -1], - [103, 102, 101, 100, 256, 257, 258, -1], - [111, 110, 384, 385, 386, -1, -1, -1], - [111, 110, 384, 385, 386, -1, -1, -1], - [111, 110, 384, 385, 386, -1, -1, -1], + [131, 130, 129, 128, 256, 257, 258, -1], + [131, 130, 129, 128, 256, 257, 258, -1], + [131, 130, 129, 128, 256, 257, 258, -1], + [257, 256, 384, 385, 386, -1, -1, -1], + [257, 256, 384, 385, 386, -1, -1, -1], + [257, 256, 384, 385, 386, -1, -1, -1], ], dtype=torch.int32, device="cuda", @@ -255,23 +246,12 @@ def test_dspark_swa_block_uses_one_scratch_page_per_request(): ) manager = DeepseekV4MemoryManager.__new__(DeepseekV4MemoryManager) - manager.full_to_swa_indexs = torch.full((32,), -1, dtype=torch.int32, device="cuda") - manager.swa_page_live_count = torch.zeros((4,), dtype=torch.int32, device="cuda") - manager._alloc_swa_pages = lambda count: torch.tensor( - [2, 0], - dtype=torch.int32, - pin_memory=True, + manager.swa_pool = SimpleNamespace(buffer=torch.empty(0, device="cuda")) + manager.swa_page_allocator = SimpleNamespace( + alloc=lambda count: torch.tensor([2, 0], dtype=torch.int32, pin_memory=True) ) mem_indexes = torch.tensor([3, 4, 5, 8, 9, 10], dtype=torch.int64, device="cuda") pages_cpu, pages = manager.alloc_dspark_swa_block(token_num=mem_indexes.numel(), block_size=3) - torch.testing.assert_close( - manager.full_to_swa_indexs[mem_indexes], - torch.full((6,), -1, dtype=torch.int32, device="cuda"), - ) - torch.testing.assert_close( - manager.swa_page_live_count, - torch.zeros((4,), dtype=torch.int32, device="cuda"), - ) torch.testing.assert_close(pages_cpu, torch.tensor([2, 0], dtype=torch.int32)) diff --git a/unit_tests/common/test_deepseek4_paged_cache.py b/unit_tests/common/test_deepseek4_paged_cache.py index bdb0b6b0f4..176f357f59 100644 --- a/unit_tests/common/test_deepseek4_paged_cache.py +++ b/unit_tests/common/test_deepseek4_paged_cache.py @@ -22,6 +22,11 @@ def cache(monkeypatch, tmp_path): json.dumps( { "page_size": 256, + "linear_att_hash_page_size": 256, + "linear_att_page_block_num": 8, + "chunked_prefill_size": 2048, + "max_req_total_len": 4096, + "disable_chunked_prefill": False, "model_dir": str(tmp_path), "penalty_counter_mode": "cpu_counter", "mtp_step": 3, @@ -101,25 +106,60 @@ def test_speculative_retry_reuses_swa_without_freeing_token_pages(cache): def cpu(data): return torch.tensor(data, dtype=torch.int32) - requests.prepare_prefill(cpu([req_idx]), cpu([0]), cpu([254]), slots[:254]) + requests.prepare_prefill(cpu([req_idx]), cpu([0]), cpu([254])) for sequences in ([255, 256, 257, 258], [255, 256, 257, 258], [256, 257, 258, 259]): - indexes = slots[cpu(sequences).cuda().long() - 1] - requests.prepare_decode(cpu([req_idx] * 4), cpu(sequences), cpu([0, 1, 2, 3]), indexes) + requests.prepare_decode(cpu([req_idx] * 4), cpu(sequences), cpu([0, 1, 2, 3])) assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages - 3 - assert manager.swa_page_live_count.sum().item() == max(sequences) + assert len(requests._swa_pages[req_idx]) == 3 assert manager.allocator.can_use_mem_size == manager.size - 512 requests.free([req_idx], slots) assert manager.allocator.can_use_mem_size == manager.size assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages - assert manager.swa_page_live_count.eq(0).all() + assert requests.req_to_swa_pages[req_idx].eq(-1).all() + + +@pytest.mark.parametrize("model_kind", ["target", "mtp", "dspark"]) +def test_model_selects_held_slots_and_prepares_only_owned_swa(cache, model_kind): + from lightllm.common.basemodel.batch_objs import ModelInput + from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + from lightllm.models.deepseek_v4_mtp.model import DeepseekV4MTPModel + from lightllm.models.deepseek_v4_dspark.model import DeepseekV4DSparkModel + + manager, requests = cache + req_idx, slots = _reserve(manager, requests, 512) + requests.prepare_swa(req_idx, 0, 256) + model_cls = {"target": DeepseekV4TpPartModel, "mtp": DeepseekV4MTPModel, "dspark": DeepseekV4DSparkModel}[ + model_kind + ] + model = model_cls.__new__(model_cls) + model.req_manager = requests + model_input = ModelInput( + batch_size=1, + total_token_num=257, + max_q_seq_len=1, + max_kv_seq_len=257, + b_req_idx=torch.tensor([req_idx], dtype=torch.int32), + b_seq_len=torch.tensor([257], dtype=torch.int32), + b_mtp_index=torch.zeros(1, dtype=torch.int32), + b_position_delta=torch.zeros(1, dtype=torch.int32), + b_shared_seq_len=torch.zeros(1, dtype=torch.int32), + b_shared_radix_node_id=torch.full((1,), -1, dtype=torch.int64), + multimodal_params=[{"images": [], "audios": []}], + ) + model_input.to_cuda() + for _ in range(2): + torch.testing.assert_close(model._select_mem_indexes(model_input), slots[256:257]) + assert len(requests._swa_pages[req_idx]) == (2 if model_kind == "dspark" else 3) + assert manager.allocator.can_use_mem_size == manager.size - 512 + requests.free([req_idx], slots) def test_prefill_chunk_keeps_all_new_swa_rows(cache): manager, requests = cache req_idx, slots = _reserve(manager, requests, 2048) - requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048]), slots) - assert manager.full_to_swa_indexs[slots].ge(0).all() - assert manager.swa_page_live_count.sum().item() == 2048 + requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048])) + assert requests.get_swa_slots(req_idx, torch.arange(2048, device="cuda")).ge(0).all() + assert len(requests._swa_pages[req_idx]) == 16 requests.free([req_idx], slots) assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages @@ -136,7 +176,7 @@ def test_history_copy_preserves_packed_data_and_scale_regions(cache): torch.testing.assert_close( pool.buffer[:, src[::256].long() // 256], pool.buffer[:, dst[::256].long() // 256], rtol=0, atol=0 ) - assert manager.full_to_swa_indexs[dst].eq(-1).all() + assert requests.req_to_swa_pages[: requests.HOLD_REQUEST_ID].eq(-1).all() def test_hold_page_and_dspark_scratch_have_separate_capacity(cache): @@ -145,7 +185,7 @@ def test_hold_page_and_dspark_scratch_have_separate_capacity(cache): torch.testing.assert_close( hold[:256], torch.arange(manager.size, manager.size + 256, dtype=torch.int32, device="cuda") ) - assert manager.full_to_swa_indexs[hold].ge(manager.swa_size).all() + assert requests.req_to_swa_pages[requests.HOLD_REQUEST_ID].eq(manager.swa_num_pages).all() assert (hold // 4).lt(manager.c4_pool.num_pages * 64).all() assert (hold // 128).lt(manager.c128_pool.num_pages * 2).all() free_tokens = manager.allocator.can_use_mem_size @@ -160,20 +200,22 @@ def test_hold_page_and_dspark_scratch_have_separate_capacity(cache): def test_cpu_cache_roundtrip_uses_derived_history_slots(cache): manager, requests = cache req_idx, src = _reserve(manager, requests, 2048) - requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048]), src) + requests.prepare_prefill(torch.tensor([req_idx]), torch.tensor([0]), torch.tensor([2048])) for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool, manager.swa_pool): pool.buffer.random_(0, 256) manager.c4_state_buffer.uniform_() manager.c4_indexer_state_buffer.uniform_() staging = torch.empty((1, manager.cpu_cache_layout.page_nbytes), dtype=torch.uint8, device="cuda") - manager.operator.pack_cpu_cache_pages(src.view(1, -1), staging) + manager.operator.pack_cpu_cache_pages(src.view(1, -1), torch.tensor([[req_idx, 2048]], device="cuda"), staging) cpu_page = torch.empty(staging.shape, dtype=torch.uint8, pin_memory=True) cpu_page.copy_(staging) - plan = manager.prepare_cpu_cache_load(2048, 2048) + dst_req = requests.alloc() + requests.prepare_swa(dst_req, 1792, 2048) + resume_slots = requests.get_swa_slots(dst_req, torch.arange(1792, 2048, device="cuda")) + plan = manager.prepare_cpu_cache_load(2048, 2048, resume_slots) manager.operator.load_cpu_cache_pages( plan, torch.tensor([0], dtype=torch.int32, device="cuda"), SimpleNamespace(cpu_kv_cache_tensor=cpu_page) ) - manager.commit_cpu_cache_load_plan(plan) for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): torch.testing.assert_close( pool.read(0, src[ratio - 1 :: ratio].long() // ratio), @@ -181,12 +223,454 @@ def test_cpu_cache_roundtrip_uses_derived_history_slots(cache): rtol=0, atol=0, ) - src_swa = manager.full_to_swa_indexs[src[-256:]].long() - dst_swa = manager.full_to_swa_indexs[plan.mem_indexes[-256:]].long() + src_swa = requests.get_swa_slots(req_idx, torch.arange(1792, 2048, device="cuda")).long() + dst_swa = resume_slots.long() for layer in range(manager.layer_num): torch.testing.assert_close( manager.swa_pool.read(layer, src_swa), manager.swa_pool.read(layer, dst_swa), rtol=0, atol=0 ) - manager.free(torch.cat([src, plan.mem_indexes])) + requests.free([req_idx, dst_req], torch.cat([src, plan.mem_indexes])) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +@pytest.mark.parametrize("length", [256, 512, 2048]) +def test_shared_history_restores_private_continuation(cache, length): + manager, requests = cache + req_idx, slots = _reserve(manager, requests, length) + requests.prepare_swa(req_idx, 0, length) + manager.swa_pool.buffer.random_(0, 256) + manager.c4_state_buffer.uniform_() + manager.c4_indexer_state_buffer.uniform_() + states = requests.create_small_page_cache_manager(1) + state_idx = states.alloc_one_state_cache() + requests.save_state(req_idx, state_idx, states, checkpoint_len=length) + torch.cuda.synchronize() + expected_swa, expected_c4, expected_indexer = [value.cuda() for value in states.get_state_cache(state_idx)] + requests.free_req(req_idx) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + reqs = [requests.alloc(), requests.alloc()] + for req in reqs: + requests.req_to_token_indexs[req, :length] = slots + requests.restore_state(SimpleNamespace(req_idx=req, cur_kv_len=0), states, state_idx, checkpoint_len=length) + pages = requests.req_to_swa_pages[req, length // 128 - 2 : length // 128].long() + torch.testing.assert_close(manager.swa_pool.buffer[:, pages], expected_swa) + tail = pages[-1] * 128 + torch.arange(124, 128, device="cuda") + rows = tail // 128 * manager.c4_state_ring + tail % manager.c4_state_ring + torch.testing.assert_close(manager.c4_state_buffer[:, rows], expected_c4) + torch.testing.assert_close(manager.c4_indexer_state_buffer[:, rows], expected_indexer) + assert set(requests._swa_pages[reqs[0]].values()).isdisjoint(requests._swa_pages[reqs[1]].values()) + # Advancing one fork releases only its own window; the shared token page survives. + requests.prepare_swa(reqs[0], length + 512, length + 768) + pages = requests.req_to_swa_pages[reqs[1], length // 128 - 2 : length // 128].long() + torch.testing.assert_close(manager.swa_pool.buffer[:, pages], expected_swa) + requests.free(reqs, slots) + states.free_state_cache([state_idx]) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + assert states.get_free_cache_num() == 1 + + +@pytest.mark.parametrize("length", [3, 127, 128, 255, 256, 257, 511, 512, 513, 2051]) +def test_pd_roundtrip_preserves_live_and_aligned_continuations(cache, length): + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4PDCacheLayout + from lightllm.models.deepseek_v4.triton_kernel.pd_cache_io import pack_pd_cache_page, unpack_pd_cache_page + + manager, requests = cache + source, src = _reserve(manager, requests, length) + destination, dst = _reserve(manager, requests, length) + requests.prepare_swa(source, 0, length) + requests.prepare_pd_decode_cache([destination], [length]) + for pool in (manager.swa_pool, manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + pool.buffer.random_(0, 256) + for buffer in (manager.c4_state_buffer, manager.c4_indexer_state_buffer, manager.c128_state_buffer): + buffer.uniform_() + layout = DeepseekV4PDCacheLayout.from_compress_rates(manager.compress_rates, token_page_size=256) + staging = torch.empty((layout.page_nbytes,), dtype=torch.uint8, device="cuda") + for start in range(0, length, 256): + end = min(length, start + 256) + pack_pd_cache_page(manager, layout, src[start:end], staging, start, length, source) + unpack_pd_cache_page(manager, layout, dst[start:end], staging, start, length, destination) + for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): + torch.testing.assert_close( + pool.read(0, src[ratio - 1 : length : ratio].long() // ratio), + pool.read(0, dst[ratio - 1 : length : ratio].long() // ratio), + rtol=0, + atol=0, + ) + checkpoint = (length - 1) // 256 * 256 + positions = torch.arange(max(0, checkpoint - 256), length, device="cuda") + src_swa = requests.get_swa_slots(source, positions).long() + dst_swa = requests.get_swa_slots(destination, positions).long() + for layer in range(manager.layer_num): + torch.testing.assert_close(manager.swa_pool.read(layer, src_swa), manager.swa_pool.read(layer, dst_swa)) + state_positions = list(range(max(0, length - 4 - length % 4), length)) + if checkpoint: + state_positions += list(range(checkpoint - 4, checkpoint)) + state_positions = torch.tensor(state_positions, device="cuda") + src_swa = requests.get_swa_slots(source, state_positions).long() + dst_swa = requests.get_swa_slots(destination, state_positions).long() + src_rows = src_swa // 128 * manager.c4_state_ring + src_swa % manager.c4_state_ring + dst_rows = dst_swa // 128 * manager.c4_state_ring + dst_swa % manager.c4_state_ring + for buffer in (manager.c4_state_buffer, manager.c4_indexer_state_buffer): + torch.testing.assert_close(buffer[:, src_rows], buffer[:, dst_rows]) + positions = torch.arange(length - length % 128, length, device="cuda") + torch.testing.assert_close( + manager.c128_state_buffer[:, source * manager.c128_state_ring + positions % manager.c128_state_ring], + manager.c128_state_buffer[:, destination * manager.c128_state_ring + positions % manager.c128_state_ring], + ) + requests.free([source, destination], torch.cat([src, dst])) assert manager.allocator.can_use_mem_size == manager.size assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +def test_dp_transfer_uses_request_page_tables(cache): + from lightllm.models.deepseek_v4.triton_kernel.dp_cache_io import copy_dsv4_dp_caches + + manager, requests = cache + source, src = _reserve(manager, requests, 512) + destination, dst = _reserve(manager, requests, 512) + requests.prepare_swa(source, 0, 512) + requests.prepare_swa(destination, 256, 512) + for pool in (manager.swa_pool, manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + pool.buffer.random_(0, 256) + manager.c4_state_buffer.uniform_() + manager.c4_indexer_state_buffer.uniform_() + pointers = torch.tensor( + [ + [ + buffer.data_ptr() + for buffer in ( + manager.c4_pool.buffer, + manager.c4_indexer_pool.buffer, + manager.c128_pool.buffer, + requests.req_to_swa_pages, + manager.swa_pool.buffer, + manager.c4_state_buffer, + manager.c4_indexer_state_buffer, + ) + ] + ], + dtype=torch.uint64, + device="cuda", + ) + meta = torch.tensor( + [0, src.data_ptr(), dst.data_ptr(), source, destination, 512], dtype=torch.uint64, device="cuda" + ) + history = torch.tensor([0, 0, 0, 1], dtype=torch.uint64, device="cuda") + copy_dsv4_dp_caches(pointers, manager, meta, history) + for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): + torch.testing.assert_close( + pool.read(0, src[ratio - 1 :: ratio].long() // ratio), pool.read(0, dst[ratio - 1 :: ratio].long() // ratio) + ) + src_swa = requests.get_swa_slots(source, torch.arange(256, 512, device="cuda")).long() + dst_swa = requests.get_swa_slots(destination, torch.arange(256, 512, device="cuda")).long() + for layer in range(manager.layer_num): + torch.testing.assert_close(manager.swa_pool.read(layer, src_swa), manager.swa_pool.read(layer, dst_swa)) + src_rows = src_swa[-4:] // 128 * manager.c4_state_ring + src_swa[-4:] % manager.c4_state_ring + dst_rows = dst_swa[-4:] // 128 * manager.c4_state_ring + dst_swa[-4:] % manager.c4_state_ring + for buffer in (manager.c4_state_buffer, manager.c4_indexer_state_buffer): + torch.testing.assert_close(buffer[:, src_rows], buffer[:, dst_rows]) + requests.free([source, destination], torch.cat([src, dst])) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +@pytest.mark.parametrize("ratio,layer", [(4, 0), (128, 1)]) +def test_compressor_checkpoint_continuation_matches_full_prefill(cache, ratio, layer): + from lightllm.models.deepseek_v4.layer_infer.compressor import fused_compress + + manager, requests = cache + source, src = _reserve(manager, requests, 512) + destination, dst = _reserve(manager, requests, 512) + requests.prepare_swa(source, 0, 512) + scores = torch.randn((512, 512 * (4 if ratio == 4 else 2)), dtype=torch.float32, device="cuda") + cos = torch.ones((512, 32), dtype=torch.float32, device="cuda") + sin = torch.zeros_like(cos) + weight = torch.ones(512, dtype=torch.float32, device="cuda") + + def compress(req_idx, slots, start, end, decode=False): + positions = torch.arange(start, end, dtype=torch.int32, device="cuda") + width = end - start + state = SimpleNamespace( + mem_manager=manager, + req_manager=requests, + mem_index=slots[start:end], + position_ids=positions, + is_prefill=not decode, + _dsv4_token_to_batch_idx=torch.zeros(width, dtype=torch.int32, device="cuda"), + b_req_idx=torch.full((width if decode else 1,), req_idx, dtype=torch.int32, device="cuda"), + b_mtp_index=torch.arange(width, dtype=torch.int32, device="cuda") + if decode + else torch.zeros(1, dtype=torch.int32, device="cuda"), + b_seq_len=positions + 1 if decode else torch.tensor([end], dtype=torch.int32, device="cuda"), + b_ready_cache_len=None if decode else torch.tensor([start], dtype=torch.int32, device="cuda"), + b_q_start_loc=None if decode else torch.zeros(1, dtype=torch.int32, device="cuda"), + dsv4_swa_write_slots=requests.get_swa_slots(req_idx, positions), + ) + fused_compress( + kv_score=scores[start:end], + infer_state=state, + layer_idx=layer, + norm_weight=weight, + eps=1e-6, + head_dim=512, + qk_rope_head_dim=64, + compress_ratio=ratio, + cos_table=cos, + sin_table=sin, + ) + + compress(source, src, 0, 512) + states = requests.create_small_page_cache_manager(1) + index = states.alloc_one_state_cache() + requests.save_state(source, index, states, checkpoint_len=256) + requests.restore_state(SimpleNamespace(req_idx=destination), states, index, checkpoint_len=256) + requests.prepare_swa(destination, 256, 512) + compress(destination, dst, 256, 508) + # Rejected speculative rows are written again in the same held token page. + compress(destination, dst, 508, 512, decode=True) + compress(destination, dst, 508, 512, decode=True) + pool = manager.c4_pool if ratio == 4 else manager.c128_pool + torch.testing.assert_close( + pool.read(0, src[256 + ratio - 1 : 512 : ratio].long() // ratio), + pool.read(0, dst[256 + ratio - 1 : 512 : ratio].long() // ratio), + rtol=0, + atol=0, + ) + requests.free([source, destination], torch.cat([src, dst])) + states.free_state_cache([index]) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + +def test_swa_index_cuda_graph_reads_updated_private_pages(cache): + from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_swa_index + + manager, requests = cache + source, src = _reserve(manager, requests, 512) + destination, dst = _reserve(manager, requests, 512) + requests.prepare_swa(source, 0, 512) + requests.prepare_swa(destination, 0, 512) + req_ids = torch.tensor([source, requests.HOLD_REQUEST_ID], dtype=torch.int32, device="cuda") + positions = torch.tensor([255, 1], dtype=torch.int32, device="cuda") + indexes = torch.empty((2, 128), dtype=torch.int32, device="cuda") + lengths = torch.empty(2, dtype=torch.int32, device="cuda") + write_slots = torch.empty_like(lengths) + + def build(): + build_swa_index(req_ids, positions, requests.req_to_swa_pages, indexes, lengths, write_slots) + + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + build() + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + build() + req_ids[0] = destination + positions[0] = 256 + graph.replay() + expected = requests.get_swa_slots(destination, torch.arange(256, 128, -1, dtype=torch.int32, device="cuda")) + torch.testing.assert_close(indexes[0], expected) + assert write_slots[0].item() == expected[0].item() + assert write_slots[1].item() == manager.swa_size + 1 + requests.free([source, destination], torch.cat([src, dst])) + + +@pytest.mark.parametrize("hit_len", [2048, 2304]) +@pytest.mark.parametrize("chunked", [True, False]) +def test_hybrid_radix_hit_fork_pause_abort_and_eviction(cache, monkeypatch, hit_len, chunked): + from sortedcontainers import SortedDict + from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context, CacheTier + from lightllm.server.router.dynamic_prompt.hybrid_att_radix_cache import HybridAttPagedRadixCache + from lightllm.utils.kv_cache_utils import compute_token_list_hash + + manager, requests = cache + small = requests.create_small_page_cache_manager(2) + radix = HybridAttPagedRadixCache(manager.size, 0, 256, 8, manager, small) + args = get_env_start_args() + args.disable_chunked_prefill = not chunked + if not chunked: + args.chunked_prefill_size = args.max_req_total_len + for name, value in { + "req_manager": requests, + "radix_cache": radix, + "args": args, + "is_hybrid_att_model": True, + "is_deepseek_v4": True, + }.items(): + monkeypatch.setattr(g_infer_context, name, value, raising=False) + + def request(req_idx, total_len): + req = InferReq.__new__(InferReq) + req.req_idx, req.args = req_idx, args + req.cur_kv_len, req.hold_kv_len, req.cur_output_len = 0, 0, 0 + req.shared_kv_node = None + req.tail_small_page_buffer_id = None + req.hybrid_len_to_big_page_id = SortedDict() + req.hybrid_cache_len = (total_len - 1) // 256 * 256 + req.image_block_spans = [] + req.cache_tiers = {CacheTier.GPU} + req.sampling_param = SimpleNamespace(disable_prompt_cache=False) + req.prompt_selected_logprobs = SimpleNamespace(copy_capture_slots_if_needed=lambda **kwargs: None) + tokens = list(range(total_len)) + hashes = compute_token_list_hash(tokens, 256) + req.shm_req = SimpleNamespace( + input_len=total_len, + shm_prompt_ids=SimpleNamespace(arr=tokens), + hybrid_token_hash_list=SimpleNamespace(size=len(hashes), get_all=lambda: hashes), + shm_cur_kv_len=0, + prompt_cache_len=0, + ) + req.get_chuncked_input_token_len = req.get_chuncked_input_token_len_for_hybrid_att + return req + + source, held = _reserve(manager, requests, 2305) + req = request(source, 2305) + req.hold_kv_len = held.numel() + for end in (2048, 2304) if chunked else (2305,): + assert req.get_chuncked_input_token_len() == end + requests.prepare_swa(source, req.cur_kv_len, end) + for pool in (manager.swa_pool, manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): + pool.buffer.random_(0, 256) + manager.c4_state_buffer.uniform_() + manager.c4_indexer_state_buffer.uniform_() + g_infer_context.save_hybrid_state_to_cache(torch.tensor([source], device="cuda"), [req]) + req.cur_kv_len = end + requests.prepare_swa(source, req.cur_kv_len, 2305) + req.cur_kv_len = 2305 + freed = [] + g_infer_context.free_a_req_mem(freed, req) + manager.free(torch.cat(freed)) + requests.free_req(source) + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + assert radix.get_tree_total_tokens_num() == 2304 + + forks = [request(requests.alloc(), hit_len + 1) for _ in range(2)] + for fork in forks: + fork._hybrid_match_radix_cache() + assert fork.cur_kv_len == hit_len + assert fork.hold_kv_len == hit_len + assert fork.shared_kv_node.node_prefix_total_len == 2048 + if hit_len == 2304: + assert requests.req_to_token_indexs[fork.req_idx, 2048].item() != held[2048].item() + restored = requests.req_to_token_indexs[fork.req_idx, :hit_len] + for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): + torch.testing.assert_close( + pool.read(0, held[ratio - 1 : hit_len : ratio].long() // ratio), + pool.read(0, restored[ratio - 1 : hit_len : ratio].long() // ratio), + ) + assert set(requests._swa_pages[forks[0].req_idx].values()).isdisjoint( + requests._swa_pages[forks[1].req_idx].values() + ) + # Pause can publish a reusable prefix; abort with GPU placement disabled only dereferences it. + forks[1].cache_tiers.clear() + for fork in forks: + freed = [] + g_infer_context.free_a_req_mem(freed, fork) + manager.free(torch.cat(freed)) + requests.free_req(fork.req_idx) + radix.free_radix_cache_to_get_enough_token(manager.size) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + assert manager.big_page_buffers.get_free_cache_num() == manager.big_page_buffers.size + assert small.get_free_cache_num() == small.size + + +def test_cpu_load_failure_releases_reserved_history_and_private_swa(cache, monkeypatch): + from lightllm.server.router.model_infer.mode_backend import dsv4_multi_level_kv_cache as module + + manager, requests = cache + req_idx = requests.alloc() + cache_module = module.Dsv4MultiLevelKvCacheModule.__new__(module.Dsv4MultiLevelKvCacheModule) + cache_module.backend = SimpleNamespace( + is_master_in_dp=False, radix_cache=None, model=SimpleNamespace(mem_manager=manager, req_manager=requests) + ) + cache_module.cpu_cache_client = SimpleNamespace() + cache_module.init_sync_group = None + req = SimpleNamespace( + req_idx=req_idx, + req_id=1, + cur_kv_len=0, + hold_kv_len=0, + image_block_spans=[], + shm_req=SimpleNamespace( + cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [0]), + token_hash_page_len_list=SimpleNamespace(get_all=lambda: [2048]), + disk_prompt_cache_len=0, + ), + ) + + def fail(**kwargs): + raise RuntimeError("injected copy failure") + + monkeypatch.setattr(manager.operator, "load_cpu_cache_pages", fail) + monkeypatch.setattr(module.g_infer_context, "get_can_alloc_token_num", lambda: manager.allocator.can_use_mem_size) + monkeypatch.setattr( + module.g_infer_context, "get_can_alloc_dsv4_swa_page_num", lambda: manager.swa_page_allocator.can_use_mem_size + ) + with pytest.raises(RuntimeError, match="injected copy failure"): + cache_module.load_cpu_cache_to_reqs([req]) + assert req.cur_kv_len == req.hold_kv_len == 0 + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + requests.free_req(req_idx) + + +def _checkpoint_transfer_worker(rank, rendezvous): + from datetime import timedelta + from sortedcontainers import SortedDict + import torch.distributed as dist + from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4CpuCacheLayout + from lightllm.common.state_cache_manager.deepseek4 import DeepseekV4StateCacheManager + from lightllm.server.router.model_infer.mode_backend.dp_backend.dp_shared_kv_trans import DPKVSharedMoudle + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + torch.cuda.set_device(rank) + dist.init_process_group( + "nccl", init_method="file://" + rendezvous, rank=rank, world_size=2, timeout=timedelta(seconds=60) + ) + try: + buffers = DeepseekV4StateCacheManager(2, DeepseekV4CpuCacheLayout.from_compress_rates([4, 128])) + if rank == 0: + for value in (31, 32): + buffers.buffer[buffers.alloc_one_state_cache()].fill_(value) + module = DPKVSharedMoudle.__new__(DPKVSharedMoudle) + module.backend = SimpleNamespace( + args=SimpleNamespace(linear_att_hash_page_size=256, linear_att_page_block_num=8, max_req_total_len=8192), + node_nccl_group=dist.group.WORLD, + model=SimpleNamespace(mem_manager=SimpleNamespace(big_page_buffers=buffers)), + radix_cache=SimpleNamespace(get_big_page_ids_by_node=lambda node: [0, 1]), + ) + for start in (0, 2048): + req = SimpleNamespace( + req_id=42, + cur_kv_len=4096 if rank == 0 else start, + shared_kv_node=object(), + hybrid_len_to_big_page_id=SortedDict(), + ) + if rank == 1 and start: + index = buffers.alloc_one_state_cache() + buffers.buffer[index].fill_(31) + req.hybrid_len_to_big_page_id[2048] = index + g_infer_context.requests_mapping = {42: req} + tasks = ( + [] + if rank == 0 + else [SimpleNamespace(max_kv_len_mem_manager_index=0, req=req, mem_indexes=range(4096 - start))] + ) + module._transfer_dsv4_checkpoints(tasks) + torch.cuda.synchronize() + if rank == 1: + assert list(req.hybrid_len_to_big_page_id) == [2048, 4096] + for length, value in ((2048, 31), (4096, 32)): + assert buffers.buffer[req.hybrid_len_to_big_page_id[length]].eq(value).all() + buffers.free_state_cache(list(req.hybrid_len_to_big_page_id.values())) + dist.barrier() + finally: + dist.destroy_process_group() + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires two CUDA devices for NCCL") +def test_dp_checkpoint_transfer_between_processes(tmp_path): + torch.multiprocessing.spawn(_checkpoint_transfer_worker, args=(str(tmp_path / "rendezvous"),), nprocs=2, join=True) diff --git a/unit_tests/models/deepseek_v4/test_vision_integration.py b/unit_tests/models/deepseek_v4/test_vision_integration.py index 80086f04d8..ad6f165306 100644 --- a/unit_tests/models/deepseek_v4/test_vision_integration.py +++ b/unit_tests/models/deepseek_v4/test_vision_integration.py @@ -116,57 +116,109 @@ class FakeRadixCache: def __init__(self): self.calls = [] - def match_prefix(self, key, update_refs=False): + def match_prefix(self, key, block_hashs, update_refs=False): self.calls.append((len(key), update_refs)) matched_len = min(len(key) // 256 * 256, 768) if update_refs: - node = SimpleNamespace(node_prefix_total_len=matched_len) + node = SimpleNamespace(node_prefix_total_len=matched_len, is_big_page_node=lambda: False) return node, matched_len, torch.arange(matched_len) return None, matched_len, None radix_cache = FakeRadixCache() - monkeypatch.setattr(g_infer_context, "is_linear_att_mixed_model", False) + monkeypatch.setattr(g_infer_context, "is_hybrid_att_model", True) + monkeypatch.setattr(g_infer_context, "get_can_alloc_dsv4_swa_page_num", lambda: 2) monkeypatch.setattr(g_infer_context, "is_deepseek_v4", True) monkeypatch.setattr(g_infer_context, "radix_cache", radix_cache) monkeypatch.setattr( g_infer_context, "req_manager", - SimpleNamespace(req_to_token_indexs=torch.empty((1, 1025), dtype=torch.int64)), + SimpleNamespace( + req_to_token_indexs=torch.empty((1, 1025), dtype=torch.int64), + restore_small_page_state=lambda **kwargs: None, + ), ) req = InferReq.__new__(InferReq) + req.args = SimpleNamespace( + page_size=256, linear_att_hash_page_size=256, linear_att_page_block_num=8, max_req_total_len=1025 + ) req.sampling_param = SimpleNamespace(disable_prompt_cache=False) req.cur_kv_len = 0 req.cur_output_len = 0 req.req_idx = 0 req.image_block_spans = [(450, 600), (700, 900)] req.shared_kv_node = None + req.tail_small_page_buffer_id = None req.shm_req = SimpleNamespace( input_len=1025, + hybrid_token_hash_list=SimpleNamespace(size=4, get_all=lambda: [1, 2, 3, 4]), shm_prompt_ids=SimpleNamespace(arr=list(range(1025))), prompt_cache_len=0, shm_cur_kv_len=0, ) - req._match_radix_cache() + req._hybrid_match_radix_cache() assert radix_cache.calls == [ (1024, False), - (700, False), - (450, False), - (450, True), + (512, False), + (256, False), + (256, True), ] assert req.cur_kv_len == 256 assert req.shm_req.prompt_cache_len == 256 +def test_checkpoint_alignment_never_lands_in_an_earlier_image(monkeypatch): + from lightllm.server.router.model_infer import infer_batch + + shm_req = SimpleNamespace( + link_prompt_ids_shm_array=lambda: None, + link_logprobs_shm_array=lambda: None, + hybrid_token_hash_list=SimpleNamespace(size=3), + ) + monkeypatch.setattr(infer_batch.g_infer_context, "is_hybrid_att_model", True) + monkeypatch.setattr( + infer_batch.g_infer_context, "shm_req_manager", SimpleNamespace(get_req_obj_by_index=lambda index: shm_req) + ) + monkeypatch.setattr( + infer_batch.g_infer_context, + "req_manager", + SimpleNamespace(req_sampling_params_manager=SimpleNamespace(init_req_sampling_params=lambda req: None)), + ) + monkeypatch.setattr( + infer_batch, + "InferSamplingParams", + lambda *args: SimpleNamespace( + pd_decode_node=None, shm_param=SimpleNamespace(stop_sequences=SimpleNamespace(to_list=lambda: [])) + ), + ) + monkeypatch.setattr(infer_batch, "PromptSelectedLogprobsExt", lambda req: None) + monkeypatch.setattr(infer_batch, "FinalTokenMetadataExt", lambda req: None) + req = infer_batch.InferReq.__new__(infer_batch.InferReq) + req.args = SimpleNamespace(linear_att_hash_page_size=256) + req.shm_index, req.vocab_size = 0, 100 + req.multimodal_params = SimpleNamespace( + to_dict=lambda: { + "images": [ + {"block_start_idx": 450, "block_end_idx": 600}, + {"block_start_idx": 700, "block_end_idx": 900}, + ] + } + ) + req._init_all_state() + assert req.hybrid_cache_len == 256 + + def test_recover_swa_budget_includes_atomic_image_block(monkeypatch): from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context monkeypatch.setattr( g_infer_context, "req_manager", - SimpleNamespace(sliding_window=128, get_prompt_cache_page_size=lambda: 256), + SimpleNamespace( + sliding_window=128, get_prompt_cache_page_size=lambda: 256, get_swa_page_need=lambda req, start, end: 16 + ), ) req = InferReq.__new__(InferReq) @@ -176,10 +228,9 @@ def test_recover_swa_budget_includes_atomic_image_block(monkeypatch): req.shm_req = SimpleNamespace(input_len=2000) req.image_block_spans = [(500, 884)] req.dsv4_swa_page_size = 128 - req.dsv4_c4_page_size = 64 - req.dsv4_has_c128 = False + req.req_idx = 0 - assert req.get_dsv4_recover_need_page_and_slot_num() == (8, 8, 0) + assert req.get_dsv4_recover_need_swa_page_num() == 8 @pytest.mark.parametrize( @@ -218,7 +269,7 @@ def get_loadable_cpu_cache_end(*args): capacity_calls.append(args) return next(capacity_results) - def prepare_cpu_cache_load(*, token_num, loaded_end): + def prepare_cpu_cache_load(*, token_num, loaded_end, resume_swa_slots): prepare_calls.append((token_num, loaded_end)) return SimpleNamespace(mem_indexes=torch.arange(token_num, dtype=torch.int32)) @@ -228,7 +279,8 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): req_to_token_indexs = torch.full((1, 6144), -1, dtype=torch.int32) req_manager = SimpleNamespace( req_to_token_indexs=req_to_token_indexs, - finish_cpu_cache_load=lambda req_idx, loaded_end: finish_calls.append((req_idx, loaded_end)), + prepare_swa=lambda req_idx, start, end: finish_calls.append((req_idx, start, end)), + get_swa_slots=lambda req_idx, positions: positions, ) mem_manager = SimpleNamespace( cpu_cache_layout=SimpleNamespace(token_page_size=2048), @@ -239,7 +291,6 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): get_loadable_cpu_cache_end=get_loadable_cpu_cache_end, prepare_cpu_cache_load=prepare_cpu_cache_load, operator=SimpleNamespace(load_cpu_cache_pages=load_cpu_cache_pages), - commit_cpu_cache_load_plan=lambda plan: None, ) module = object.__new__(cache_module.Dsv4MultiLevelKvCacheModule) module.backend = SimpleNamespace( @@ -276,8 +327,8 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): monkeypatch.setattr(cache_module.g_infer_context, "get_can_alloc_token_num", lambda: 8192) monkeypatch.setattr( cache_module.g_infer_context, - "get_can_alloc_dsv4_page_and_slot_num", - lambda: (2, 0, 0), + "get_can_alloc_dsv4_swa_page_num", + lambda: 2, ) module.load_cpu_cache_to_reqs([req]) @@ -285,7 +336,7 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): assert len(capacity_calls) == 2 assert prepare_calls == [(2048, 2048)] assert loaded_pages == [[10]] - assert finish_calls == [(0, 2048)] + assert finish_calls == [(0, 1792, 2048)] assert req.cur_kv_len == 2048 assert req.shm_req.shm_cur_kv_len == 2048 assert req.shm_req.cpu_prompt_cache_len == 2048 @@ -401,8 +452,8 @@ def test_swa_index_adds_bidirectional_image_visibility(): positions = torch.tensor([5, 8, 10], dtype=torch.int32, device="cuda") req_idx = torch.zeros(3, dtype=torch.int32, device="cuda") - req_to_token = torch.arange(20, dtype=torch.int32, device="cuda").unsqueeze(0) - full_to_swa = torch.arange(100, 120, dtype=torch.int32, device="cuda") + req_to_swa_pages = torch.tensor([[1]], dtype=torch.int32, device="cuda") + write_slots = torch.empty_like(positions) output = torch.empty((3, 8), dtype=torch.int32, device="cuda") lengths = torch.empty(3, dtype=torch.int32, device="cuda") image_left = torch.tensor([0, 2, 4], dtype=torch.int32, device="cuda") @@ -411,18 +462,18 @@ def test_swa_index_adds_bidirectional_image_visibility(): build_swa_index( req_idx, positions, - req_to_token, - full_to_swa, + req_to_swa_pages, output, lengths, + write_slots, window=4, image_left=image_left, image_right=image_right, ) assert output.cpu().tolist() == [ - [105, 104, 103, 102, -1, -1, -1, -1], - [105, 106, 107, 108, 109, 110, -1, -1], - [106, 107, 108, 109, 110, -1, -1, -1], + [133, 132, 131, 130, -1, -1, -1, -1], + [133, 134, 135, 136, 137, 138, -1, -1], + [134, 135, 136, 137, 138, -1, -1, -1], ] assert lengths.cpu().tolist() == [4, 6, 5] diff --git a/unit_tests/server/router/model_infer/mode_backend/test_paged_side_allocations.py b/unit_tests/server/router/model_infer/mode_backend/test_paged_side_allocations.py index 45fc0017c2..1e0f98383c 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_paged_side_allocations.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_paged_side_allocations.py @@ -39,6 +39,7 @@ def test_pd_decode_reserves_model_pages_but_transfers_only_logical_kv(monkeypatc next_index = [4] tasks = [] backend = decode_impl.PDDecodeNode.__new__(decode_impl.PDDecodeNode) + backend.is_deepseek_v4 = False backend.args = SimpleNamespace(pd_kv_page_size=4, page_size=4) backend.model = SimpleNamespace(req_manager=SimpleNamespace(req_to_token_indexs=table)) backend.is_master_in_dp = False @@ -85,6 +86,7 @@ def create_task(**kwargs): def test_pd_decode_rejects_unaligned_transfer_page_size(): backend = decode_impl.PDDecodeNode.__new__(decode_impl.PDDecodeNode) + backend.is_deepseek_v4 = False backend.args = SimpleNamespace(pd_kv_page_size=3, page_size=4) req = SimpleNamespace(cur_kv_len=4, shm_req=SimpleNamespace(input_len=10)) diff --git a/unit_tests/server/test_pd_start_args.py b/unit_tests/server/test_pd_start_args.py index def9608eb8..deadbd8fe0 100644 --- a/unit_tests/server/test_pd_start_args.py +++ b/unit_tests/server/test_pd_start_args.py @@ -4,6 +4,49 @@ from lightllm.server.core.objs.start_args_type import StartArgs +@pytest.mark.parametrize( + "page_size,hash_page_size,cpu_page_size", [(1, 256, None), (256, 512, None), (256, 256, 4096), (256, 256, None)] +) +def test_dsv4_page_and_checkpoint_configuration(monkeypatch, page_size, hash_page_size, cpu_page_size): + for name in ( + "_set_envs_and_config", + "auto_set_max_req_total_len", + "auto_set_fused_shared_experts", + "set_unique_server_name", + ): + monkeypatch.setattr("lightllm.server.api_start." + name, lambda args: None) + monkeypatch.setattr("lightllm.server.api_start.get_model_type", lambda model_dir: "deepseek_v4") + monkeypatch.setattr("lightllm.server.api_start.is_hybrid_att_model", lambda model_dir: True) + + def validation_finished(args): + assert args.cpu_cache_token_page_size == 2048 + raise RuntimeError("DSV4 page-size validation passed") + + monkeypatch.setattr("lightllm.server.api_start.auto_set_response_parsers", validation_finished) + args = StartArgs( + model_dir="unused", + page_size=page_size, + linear_att_hash_page_size=hash_page_size, + linear_att_page_block_num=8, + enable_cpu_cache=True, + cpu_cache_token_page_size=cpu_page_size, + max_req_total_len=8192, + eos_id=[2], + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + if page_size != 256 or hash_page_size != 256: + with pytest.raises(ValueError, match="DeepSeek-V4 requires"): + _launch_subprocesses(args) + elif cpu_page_size is not None: + with pytest.raises(ValueError, match="CPU cache pages must match"): + _launch_subprocesses(args) + else: + with pytest.raises(RuntimeError, match="DSV4 page-size validation passed"): + _launch_subprocesses(args) + + def test_pd_kv_page_size_must_be_divisible_by_model_page_size(monkeypatch): monkeypatch.setattr("lightllm.server.api_start._set_envs_and_config", lambda args: None) monkeypatch.setattr("lightllm.server.api_start.auto_set_max_req_total_len", lambda args: None) From 889abf514d79a75dbccb3ceff87669f0f3851601 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 23:03:35 +0800 Subject: [PATCH 193/214] Use V4 SWA slots for DSpark draft KV writes --- .../deepseek_v4_dspark/layer_infer/transformer_layer_infer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py index d8d5a30ffe..0c075073f4 100644 --- a/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4_dspark/layer_infer/transformer_layer_infer.py @@ -31,7 +31,7 @@ def context_forward( infer_state.context_kv = all_kv.split(self.head_dim_, dim=-1) infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope( layer_index=self.layer_num_, - mem_index=infer_state.mem_index, + swa_slots=infer_state.dsv4_swa_write_slots, kv=infer_state.context_kv[self.stage_id], kv_weight=layer_weight.kv_norm_.weight, eps=self.eps_, From 7ab3a6660c79198f6b152d96352eb36f2d6beb87 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Sun, 20 Sep 2026 23:34:56 +0800 Subject: [PATCH 194/214] Initialize DSpark commit SWA metadata in decode --- lightllm/models/deepseek_v4_dspark/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/models/deepseek_v4_dspark/model.py b/lightllm/models/deepseek_v4_dspark/model.py index 40866be713..6dca7ebc84 100644 --- a/lightllm/models/deepseek_v4_dspark/model.py +++ b/lightllm/models/deepseek_v4_dspark/model.py @@ -205,7 +205,7 @@ def _decode(self, model_input: ModelInput) -> ModelOutput: assert model_input.mtp_draft_input_hiddens.shape[0] == model_input.batch_size infer_state = self._create_inferstate(model_input) - infer_state.position_ids = model_input.b_seq_len - 1 + infer_state.init_some_extra_state(self) 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): From 71d2ae99bfa8097145ee3e3ce96f07a27e451cb2 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 21 Sep 2026 00:07:12 +0800 Subject: [PATCH 195/214] Allow read-only hybrid radix prefix probes --- .../dynamic_prompt/hybrid_att_radix_cache.py | 22 +++++---- .../router/dynamic_prompt/test_radix_cache.py | 48 +++++++++++++++++++ 2 files changed, 60 insertions(+), 10 deletions(-) diff --git a/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py b/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py index f1cff5fe87..5f95e4aba5 100644 --- a/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py +++ b/lightllm/server/router/dynamic_prompt/hybrid_att_radix_cache.py @@ -342,7 +342,6 @@ def match_prefix( block_hashs: Optional[List[int]] = None, update_refs: bool = False, ): - assert update_refs is True, "update_refs must be True" assert key is not None, "key must not be None" if block_hashs is None: block_hashs = [] @@ -366,7 +365,7 @@ def match_prefix( return None, 0, None # 判定真正可以用的匹配节点。 - ans_node_list = self._trim_unusable_match_tail(ans_node_list) + ans_node_list = self._trim_unusable_match_tail(ans_node_list, update_refs=update_refs) if len(ans_node_list) == 0: return None, 0, None @@ -431,7 +430,9 @@ def _match_prefix_helper( finally: self._add_node(node) - def _trim_unusable_match_tail(self, nodes: List[HybridAttPagedTreeNode]) -> List[HybridAttPagedTreeNode]: + def _trim_unusable_match_tail( + self, nodes: List[HybridAttPagedTreeNode], update_refs: bool + ) -> List[HybridAttPagedTreeNode]: removed_list = [] for node in reversed(nodes): if node.is_big_page_node(): @@ -442,14 +443,15 @@ def _trim_unusable_match_tail(self, nodes: List[HybridAttPagedTreeNode]) -> List else: removed_list.append(node) - for node in removed_list: - self._discard_node(node) - # dec ref - node.ref_counter -= 1 - if node.ref_counter == 0: - self.refed_tokens_num.arr[0] -= len(node.token_mem_index_value) + if update_refs: + for node in removed_list: + self._discard_node(node) + # dec ref + node.ref_counter -= 1 + if node.ref_counter == 0: + self.refed_tokens_num.arr[0] -= len(node.token_mem_index_value) - self._add_node(node) + self._add_node(node) if len(removed_list) == 0: return nodes diff --git a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py index 7611734bd1..9908b9d49c 100644 --- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py +++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py @@ -2,7 +2,9 @@ import pytest import torch +from sortedcontainers import SortedDict +from lightllm.server.router.dynamic_prompt.hybrid_att_radix_cache import HybridAttPagedRadixCache from lightllm.server.router.dynamic_prompt.radix_cache import RadixCache, TreeNode from lightllm.utils import shm_utils @@ -324,5 +326,51 @@ def test_page_key_bytes_does_not_share_tensor_memory(): assert key != TreeNode(page_size=4).get_child_key(token_ids) +def test_hybrid_dry_match_keeps_refs_after_tail_checkpoint_eviction(): + freed_state_ids = [] + small_page_buffers = SimpleNamespace( + get_free_cache_num=lambda: 0, + free_state_cache=lambda free_indexes: freed_state_ids.extend(free_indexes), + ) + mem_manager = SimpleNamespace(big_page_buffers=SimpleNamespace()) + tree = HybridAttPagedRadixCache( + total_token_num=100, + rank_in_node=101, + hash_page_size=4, + big_page_num=2, + kv_cache_mem_manager=mem_manager, + small_page_buffers=small_page_buffers, + ) + key = torch.arange(12, dtype=torch.int64) + values = torch.arange(100, 112, dtype=torch.int64) + tree.insert( + key, + values, + block_hashs=[11, 22, 33], + block_state_idxs=[None, None, 7], + len_to_big_page_id=SortedDict({8: 5}), + ) + + big_page = tree.root_node.children[22] + tail_page = big_page.children[33] + node, matched_len, _ = tree.match_prefix(key, block_hashs=[11, 22, 33], update_refs=False) + assert node is tail_page and matched_len == 12 + assert tree.get_refed_tokens_num() == 0 + + tree.free_one_small_page_buffer() + assert freed_state_ids == [7] + node, matched_len, matched_values = tree.match_prefix(key, block_hashs=[11, 22, 33], update_refs=False) + assert node is big_page and matched_len == 8 + assert matched_values.tolist() == list(range(100, 108)) + assert big_page.ref_counter == tail_page.ref_counter == tree.get_refed_tokens_num() == 0 + + node, matched_len, _ = tree.match_prefix(key, block_hashs=[11, 22, 33], update_refs=True) + assert node is big_page and matched_len == 8 + assert big_page.ref_counter == 1 and tail_page.ref_counter == 0 + assert tree.get_refed_tokens_num() == 8 + tree.dec_node_ref_counter(node) + assert tree.get_refed_tokens_num() == 0 + + if __name__ == "__main__": pytest.main() From 00fe6509bbe35429c2cc399fd9ff053981321a03 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 21 Sep 2026 09:13:22 +0800 Subject: [PATCH 196/214] Reuse GPU staging for V4 DP checkpoint transfers --- .../mode_backend/dp_backend/dp_shared_kv_trans.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index ef9cf5c3f2..50655ac7d7 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -170,6 +170,7 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): all_tasks = [None for _ in range(dist.get_world_size(group=group))] dist.all_gather_object(all_tasks, local_tasks, group=group) buffers = self.backend.model.mem_manager.big_page_buffers + staging = None for destination, tasks in enumerate(all_tasks): for source, req_id, start, end in tasks: first = (start // big_tokens + 1) * big_tokens @@ -177,6 +178,8 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): if not lengths or rank not in (source, destination): continue req = g_infer_context.requests_mapping[req_id] + if staging is None: + staging = torch.empty((buffers.buffer.shape[1],), dtype=torch.uint8, device="cuda") if rank == source: shared_ids = self.backend.radix_cache.get_big_page_ids_by_node(req.shared_kv_node) indexes = [ @@ -186,10 +189,9 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): for length in lengths ] for index in indexes: - staging = buffers.buffer[index].cuda(non_blocking=True) + staging.copy_(buffers.buffer[index], non_blocking=True) dist.send(staging, dst=dist.get_global_rank(group, destination), group=group) else: - staging = torch.empty((buffers.buffer.shape[1],), dtype=torch.uint8, device="cuda") for length in lengths: index = buffers.alloc_one_state_cache() assert index is not None From 6ca777960f4aa6ba946c47920f62912663cfd0fb Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 21 Sep 2026 09:22:22 +0800 Subject: [PATCH 197/214] Use shared NCCL group for V4 checkpoint P2P --- .../mode_backend/dp_backend/dp_shared_kv_trans.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index 50655ac7d7..ea15b3af6c 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -190,12 +190,14 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): ] for index in indexes: staging.copy_(buffers.buffer[index], non_blocking=True) - dist.send(staging, dst=dist.get_global_rank(group, destination), group=group) + send_op = dist.P2POp(dist.isend, staging, group=group, group_peer=destination) + dist.batch_isend_irecv([send_op])[0].wait() else: for length in lengths: index = buffers.alloc_one_state_cache() assert index is not None - dist.recv(staging, src=dist.get_global_rank(group, source), group=group) + recv_op = dist.P2POp(dist.irecv, staging, group=group, group_peer=source) + dist.batch_isend_irecv([recv_op])[0].wait() buffers.buffer[index].copy_(staging, non_blocking=True) req.hybrid_len_to_big_page_id[length] = index From 1fda8b77cfe2aba66d74f694693e85eac8b3da29 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 21 Sep 2026 15:11:53 +0800 Subject: [PATCH 198/214] Transfer V4 DP checkpoints through CPU process group --- .../dp_backend/dp_shared_kv_trans.py | 21 +++++++------------ 1 file changed, 8 insertions(+), 13 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index ea15b3af6c..86efda9586 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -148,15 +148,15 @@ def build_shared_kv_trans_tasks( def _transfer_dsv4_checkpoints(self, trans_tasks): """Transfer CPU snapshots for the history fetched from another DP rank. - CUDA IPC exports the packed history and private runtime only. NCCL moves - the checkpoint bytes through temporary GPU tensors so pinned host pools - keep their process-local ownership and registration. + CUDA IPC exports the packed history and private runtime only. The + checkpoint pools are process-local pinned CPU memory, so move their + bytes through the node's CPU process group. """ args = self.backend.args big_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num if big_tokens > args.max_req_total_len: return - group = self.backend.node_nccl_group + group = self.backend.node_gloo_group rank = dist.get_rank(group=group) local_tasks = [ ( @@ -170,7 +170,6 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): all_tasks = [None for _ in range(dist.get_world_size(group=group))] dist.all_gather_object(all_tasks, local_tasks, group=group) buffers = self.backend.model.mem_manager.big_page_buffers - staging = None for destination, tasks in enumerate(all_tasks): for source, req_id, start, end in tasks: first = (start // big_tokens + 1) * big_tokens @@ -178,9 +177,8 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): if not lengths or rank not in (source, destination): continue req = g_infer_context.requests_mapping[req_id] - if staging is None: - staging = torch.empty((buffers.buffer.shape[1],), dtype=torch.uint8, device="cuda") if rank == source: + peer = dist.get_global_rank(group, destination) shared_ids = self.backend.radix_cache.get_big_page_ids_by_node(req.shared_kv_node) indexes = [ req.hybrid_len_to_big_page_id[length] @@ -189,16 +187,13 @@ def _transfer_dsv4_checkpoints(self, trans_tasks): for length in lengths ] for index in indexes: - staging.copy_(buffers.buffer[index], non_blocking=True) - send_op = dist.P2POp(dist.isend, staging, group=group, group_peer=destination) - dist.batch_isend_irecv([send_op])[0].wait() + dist.send(buffers.buffer[index], dst=peer, group=group) else: + peer = dist.get_global_rank(group, source) for length in lengths: index = buffers.alloc_one_state_cache() assert index is not None - recv_op = dist.P2POp(dist.irecv, staging, group=group, group_peer=source) - dist.batch_isend_irecv([recv_op])[0].wait() - buffers.buffer[index].copy_(staging, non_blocking=True) + dist.recv(buffers.buffer[index], src=peer, group=group) req.hybrid_len_to_big_page_id[length] = index def kv_trans(self, trans_tasks: List["TransTask"]): From 963d3481416cc1fea515afc79fc9174ac92c1787 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 22 Sep 2026 12:23:17 +0800 Subject: [PATCH 199/214] docs(test): align PD defaults after main merge --- docs/CN/source/tutorial/api_server_args.rst | 10 +++++----- docs/EN/source/tutorial/api_server_args.rst | 11 ++++++----- unit_tests/utils/test_envs_utils.py | 16 ++++++++-------- 3 files changed, 19 insertions(+), 18 deletions(-) diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index 53027e97a6..74287bcc04 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -139,7 +139,7 @@ PD 分离模式参数 ``pd_node_resource_wait_timeout_seconds`` 为所有请求下发统一的资源等待上限;P/D 节点只负责按下发值 控制本地 ``shm_req`` 申请和 Router 等待进入推理系统,不读取本地限流开关或超时配置。首段的等待上限由 PD Master 上的 - ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 10 秒;设置为 -1 表示永久等待。 + ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 20 秒;设置为 -1 表示永久等待。 ``segment_index > 0`` 的续跑分段使用独立的等待上限,该值由 ``LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 60 秒,以提高已产生部分结果的 请求最终完成的成功率。 @@ -150,15 +150,15 @@ PD 分离模式参数 则不再从头重试,以免产生重复内容。设置 ``--disable_pd_node_self_request_limit`` 后,PD Master 不再下发 有限的资源等待时间;P/D 节点永久等待,其他原因产生的 ``Server is busy`` 也会直接返回,不触发重试。 多机 TP 场景仅由 master 节点执行超时判断,slave 节点永久等待。cache 命中记录允许提升优先级的最大年龄由 - ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` 控制,默认 36 秒。cache 命中提权还要求输入 - token 数达到 ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` 配置的门槛(默认 4096),避免短请求仅因 + ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` 控制,默认 180 秒。cache 命中提权还要求输入 + token 数达到 ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` 配置的门槛(默认 2048),避免短请求仅因 cache 命中率高而提升优先级。 启动示例: .. code-block:: bash - LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=10 \ + LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=20 \ LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS=60 \ LIGHTLLM_PD_NODE_BUSY_RETRY_TIMEOUT_SECONDS=120 \ python -m lightllm.server.api_server --run_mode pd_master ... @@ -238,7 +238,7 @@ PD 分离模式参数 .. option:: --running_max_req_size - 同时进行前向推理的最大请求数量,默认为 ``1000`` + 同时进行前向推理的最大请求数量,默认为 ``256``。PD Master 当前不使用该参数执行请求准入限流。 .. option:: --max_req_total_len diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 0f82294c61..0085e3737b 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -147,7 +147,7 @@ PD disaggregation Mode Parameters ``pd_node_resource_wait_timeout_seconds`` for every request. P/D nodes only enforce the received value for local ``shm_req`` allocation and the wait from Router entry to inference entry; they do not read local limiting switches or timeout settings. The first segment's timeout is - controlled on PD Master by ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 10 seconds; set it + controlled on PD Master by ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 20 seconds; set it to -1 to wait indefinitely. Continuation segments with ``segment_index > 0`` use a separate timeout controlled by ``LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 60 seconds, improving the chance that requests which have already produced partial results complete successfully. When set to a non-negative value, @@ -162,15 +162,15 @@ PD disaggregation Mode Parameters In multi-node TP deployments, only the master node evaluates the timeout; slave nodes wait indefinitely. The maximum cache-record age eligible for promotion is controlled by ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` and defaults to - 36 seconds. Cache-hit promotion also requires at least the number of input tokens configured by - ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` (4096 by default), so short requests do not gain priority + 180 seconds. Cache-hit promotion also requires at least the number of input tokens configured by + ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` (2048 by default), so short requests do not gain priority solely from a high cache-hit rate. Startup example: .. code-block:: bash - LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=10 \ + LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=20 \ LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS=60 \ LIGHTLLM_PD_NODE_BUSY_RETRY_TIMEOUT_SECONDS=120 \ python -m lightllm.server.api_server --run_mode pd_master ... @@ -253,7 +253,8 @@ Memory and Batch Processing Parameters .. option:: --running_max_req_size - Maximum number of requests for simultaneous forward inference, default is ``1000`` + Maximum number of requests for simultaneous forward inference, default ``256``. + PD Master does not currently use this parameter for request admission limiting. .. option:: --max_req_total_len diff --git a/unit_tests/utils/test_envs_utils.py b/unit_tests/utils/test_envs_utils.py index 0cac7caa9e..f4060e6287 100644 --- a/unit_tests/utils/test_envs_utils.py +++ b/unit_tests/utils/test_envs_utils.py @@ -7,11 +7,11 @@ ) -def test_pd_cache_high_priority_max_age_defaults_to_36_seconds(monkeypatch): +def test_pd_cache_high_priority_max_age_defaults_to_180_seconds(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS", raising=False) get_pd_cache_high_priority_max_age_seconds.cache_clear() - assert get_pd_cache_high_priority_max_age_seconds() == 36 + assert get_pd_cache_high_priority_max_age_seconds() == 180 get_pd_cache_high_priority_max_age_seconds.cache_clear() @@ -25,29 +25,29 @@ def test_pd_cache_high_priority_max_age_reads_environment_variable(monkeypatch): get_pd_cache_high_priority_max_age_seconds.cache_clear() -def test_pd_cache_high_priority_min_prompt_tokens_defaults_to_4096(monkeypatch): +def test_pd_cache_high_priority_min_prompt_tokens_defaults_to_2048(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", raising=False) get_pd_cache_high_priority_min_prompt_tokens.cache_clear() - assert get_pd_cache_high_priority_min_prompt_tokens() == 4096 + assert get_pd_cache_high_priority_min_prompt_tokens() == 2048 get_pd_cache_high_priority_min_prompt_tokens.cache_clear() def test_pd_cache_high_priority_min_prompt_tokens_reads_environment_variable(monkeypatch): - monkeypatch.setenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", "2048") + monkeypatch.setenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", "4096") get_pd_cache_high_priority_min_prompt_tokens.cache_clear() - assert get_pd_cache_high_priority_min_prompt_tokens() == 2048 + assert get_pd_cache_high_priority_min_prompt_tokens() == 4096 get_pd_cache_high_priority_min_prompt_tokens.cache_clear() -def test_pd_node_resource_wait_timeout_defaults_to_10_seconds(monkeypatch): +def test_pd_node_resource_wait_timeout_defaults_to_20_seconds(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS", raising=False) get_pd_node_resource_wait_timeout_seconds.cache_clear() - assert get_pd_node_resource_wait_timeout_seconds() == 10 + assert get_pd_node_resource_wait_timeout_seconds() == 20 get_pd_node_resource_wait_timeout_seconds.cache_clear() From 25f605ffcbe2e54912d95baaaf892bceb797ad0d Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 22 Sep 2026 05:07:25 +0000 Subject: [PATCH 200/214] dsv4 default page size 256 --- lightllm/server/api_start.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index abcb074a7d..489d27cfa6 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -75,6 +75,16 @@ def _launch_subprocesses(args: StartArgs): if args.run_mode not in ["normal", "prefill", "decode", "visual_only"]: return + if model_type == "deepseek_v4": + if args.page_size != 256 or args.linear_att_hash_page_size != 256: + logger.warning( + "DeepSeek-V4 forces --page_size and --linear_att_hash_page_size to 256 (got %s and %s)", + args.page_size, + args.linear_att_hash_page_size, + ) + args.page_size = 256 + args.linear_att_hash_page_size = 256 + # 通过模型的参数判断是否是多模态模型,包含哪几种模态, 并设置是否启动相应得模块 if args.disable_vision is None: if has_vision_module(args.model_dir): @@ -215,9 +225,6 @@ def _launch_subprocesses(args: StartArgs): if args.page_size < 1: raise ValueError(f"--page_size must be >= 1, got {args.page_size}") - if get_model_type(args.model_dir) == "deepseek_v4": - if args.page_size != 256 or args.linear_att_hash_page_size != 256: - raise ValueError("DeepSeek-V4 requires --page_size 256 --linear_att_hash_page_size 256") if args.run_mode in ("prefill", "decode"): assert args.pd_kv_page_size % args.page_size == 0, "--pd_kv_page_size must be divisible by --page_size" From e2593c1f9554b518da82d23fdfd5b1aa9bbcf0a8 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 22 Sep 2026 09:56:48 +0000 Subject: [PATCH 201/214] fix(ds4): size decode SWA by request window and cap startup warmup --- lightllm/common/basemodel/basemodel.py | 15 ++++-- .../deepseek4_mem_manager.py | 51 ++++++++++--------- lightllm/common/req_manager/deepseek4.py | 4 +- lightllm/models/deepseek_v4/model.py | 18 ++++++- 4 files changed, 57 insertions(+), 31 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 2e0b743a34..d087e24124 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -303,9 +303,10 @@ def _init_cudagraph(self): self.graph.warmup(self) def _init_prefill_cuda_graph(self): + # Draft models use self.run_mode="normal" even when the node is decode-only. self.prefill_graph = ( None - if not get_env_start_args().enable_prefill_cudagraph + if self.args.run_mode == "decode" or not get_env_start_args().enable_prefill_cudagraph else PrefillCudaGraph(decode_cuda_graph=self.graph, tp_world_size=self.tp_world_size_) ) if self.prefill_graph is not None: @@ -988,6 +989,8 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: @final @torch.no_grad() def _check_max_len_infer(self): + if self.args.run_mode == "decode": + return disable_check_max_len_infer = os.getenv("DISABLE_CHECK_MAX_LEN_INFER", None) is not None if disable_check_max_len_infer: logger.info("disable_check_max_len_infer is true") @@ -1063,12 +1066,16 @@ def _autotune_warmup(self): Autotuner.start_autotune_warmup(AutotuneKernelType.GENERAL) torch.distributed.barrier() + warmup_max_tokens = self.batch_max_tokens + if self.args.run_mode == "decode": + decode_rows = self.max_req_num * self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + warmup_max_tokens = min(warmup_max_tokens, decode_rows) warmup_lengths = [1, 4, 8, 16, 32, 64, 128, 256, 1024, 2048, 4096] - if self.batch_max_tokens not in warmup_lengths: - warmup_lengths.append(self.batch_max_tokens) + if warmup_max_tokens not in warmup_lengths: + warmup_lengths.append(warmup_max_tokens) - warmup_lengths = [e for e in warmup_lengths if e <= self.batch_max_tokens] + warmup_lengths = [e for e in warmup_lengths if e <= warmup_max_tokens] warmup_lengths.sort(reverse=True) diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index ea3aae2a5f..6b306714a6 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -1,5 +1,3 @@ -import os - import torch from dataclasses import dataclass from typing import List, Optional, Sequence @@ -38,9 +36,6 @@ # 追加候选槽,避免 rejected draft 覆盖仍存活的基础窗口。c128 同样追加候选槽,再对齐到 ratio 4。 DSV4_C4_STATE_RING = 8 # 8 rows/page before MTP padding DSV4_C128_STATE_RING = 128 # 128 rows/request before MTP padding -# swa 池占 full token 空间的比例(sglang DSV4 默认 swa_full_tokens_ratio=0.1 同值)。 -# 瞬时借页/驱逐走 swa 压力阀;池子大小仅按 ratio 切分,不再叠加结构性余量。 -DSV4_SWA_FULL_TOKENS_RATIO = float(os.getenv("DSV4_SWA_FULL_TOKENS_RATIO", "0.1")) def _ceil_div(a: int, b: int) -> int: @@ -425,15 +420,13 @@ def __init__( compress_rates: List[int], max_request_num: int, mtp_step: int, + swa_page_num: int, indexer_head_dim: int = 128, cpu_cache_token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, - swa_full_tokens_ratio: float = DSV4_SWA_FULL_TOKENS_RATIO, always_copy=False, mem_fraction=0.9, memory_reservations=None, ): - if get_env_start_args().page_size != DSV4_PROMPT_CACHE_PAGE_SIZE: - raise ValueError("DeepSeek-V4 requires --page_size 256") assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" assert head_dim == self.mla_head_dim, f"DeepSeek-V4 packed KV 期望 head_dim={self.mla_head_dim}" assert ( @@ -443,6 +436,7 @@ def __init__( assert all(r in (0, 4, 128) for r in compress_rates), "compress_rates 取值只能是 0/4/128" assert max_request_num > 0, "max_request_num 必须为正数" assert 0 <= mtp_step < DSV4_C128_STATE_RING, "mtp_step 必须位于 [0, 128)" + assert swa_page_num > 0, "swa_page_num 必须为正数" self.compress_rates = list(compress_rates) self.n_c4 = sum(1 for r in self.compress_rates if r == 4) @@ -451,7 +445,7 @@ def __init__( self.max_request_num = max_request_num self.c4_state_ring = DSV4_C4_STATE_RING + mtp_step self.c128_state_ring = _ceil_div(DSV4_C128_STATE_RING + mtp_step, 4) * 4 - self.swa_full_tokens_ratio = float(swa_full_tokens_ratio) + self.swa_page_num = int(swa_page_num) self.cpu_cache_layout = DeepseekV4CpuCacheLayout.from_compress_rates( self.compress_rates, token_page_size=cpu_cache_token_page_size, @@ -483,9 +477,6 @@ def __init__( ) # ------------------------------------------------------------------ sizing - def _planned_swa_size(self, full_size: int) -> int: - return _ceil_div(int(full_size * self.swa_full_tokens_ratio), DSV4_SWA_PAGE_SIZE) * DSV4_SWA_PAGE_SIZE - @staticmethod def _paged_state_rows(num_swa_pages: int, ring: int, ratio: int) -> int: rows = num_swa_pages * ring + ring + 1 @@ -499,19 +490,31 @@ def _init_state_sentinel(buffer: torch.Tensor) -> None: return def get_cell_size(self): - kv_bytes = self.mla_bytes_per_token - indexer_bytes = self.indexer_bytes_per_token - state_dtype_bytes = torch._utils._element_size(torch.float32) - c4_state_width = 4 * self.head_dim + 4 * self.indexer_head_dim - c4_state_bytes = self.c4_state_ring / DSV4_SWA_PAGE_SIZE * c4_state_width * state_dtype_bytes * self.n_c4 - swa_slot = kv_bytes * self.layer_num + c4_state_bytes - compressed = (kv_bytes + indexer_bytes) * self.n_c4 / 4 + kv_bytes * self.n_c128 / 128 - - return swa_slot * self.swa_full_tokens_ratio + compressed + # Profile the physical PackedPagePool stride, not 584 logical bytes per MLA slot. + # Its MLA pages round (slots * (576 data + 8 scale)) up to a 576-byte + # data-row boundary; the indexer pool uses align_bytes=1 (no padding). + # One 256-token full page owns one C4/indexer page and one C128 page + # per corresponding layer, so divide their total bytes by 256. + c4_page = _aligned_gpu_page_nbytes( + DSV4_C4_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + indexer_page = _aligned_gpu_page_nbytes(DSV4_C4_PAGE_SIZE, self.indexer_head_dim, DSV4_INDEXER_SCALE_BYTES) + c128_page = _aligned_gpu_page_nbytes( + DSV4_C128_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + return (self.n_c4 * (c4_page + indexer_page) + self.n_c128 * c128_page) / DSV4_PROMPT_CACHE_PAGE_SIZE def get_fixed_memory_size(self): - state_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 - return self.n_c128 * state_rows * (2 * self.head_dim) * torch._utils._element_size(torch.float32) + swa_page = _aligned_gpu_page_nbytes( + DSV4_SWA_PAGE_SIZE, DSV4_MLA_DATA_BYTES_PER_TOKEN, self.mla_scale_bytes, DSV4_MLA_PAGE_ALIGN_BYTES + ) + swa_bytes = self.layer_num * (self.swa_page_num + 1) * swa_page # includes the HOLD page + c4_rows = self._paged_state_rows(self.swa_page_num, self.c4_state_ring, 4) + c4_state_bytes = self.n_c4 * c4_rows * (4 * self.head_dim + 4 * self.indexer_head_dim) * 4 + c128_rows = (self.max_request_num + 1) * self.c128_state_ring + 1 + c128_state_bytes = self.n_c128 * c128_rows * (2 * self.head_dim) * 4 + compressed_hold_bytes = int(self.get_cell_size() * DSV4_PROMPT_CACHE_PAGE_SIZE) + return swa_bytes + c4_state_bytes + c128_state_bytes + compressed_hold_bytes def get_pd_kv_move_buffer_size(self): args = get_env_start_args() @@ -530,7 +533,7 @@ def _init_buffers(self, size, dtype, head_num, head_dim, layer_num): rank_in_node = get_current_rank_in_node() server = get_unique_server_name() - self.swa_size = self._planned_swa_size(size) + self.swa_size = self.swa_page_num * DSV4_SWA_PAGE_SIZE self.swa_pool = PackedPagePool( size=self.swa_size, page_size=DSV4_SWA_PAGE_SIZE, diff --git a/lightllm/common/req_manager/deepseek4.py b/lightllm/common/req_manager/deepseek4.py index 58bcfaa6b5..29d612e4df 100644 --- a/lightllm/common/req_manager/deepseek4.py +++ b/lightllm/common/req_manager/deepseek4.py @@ -63,7 +63,7 @@ def prepare_swa(self, req_idx, start, end): first_retained_page = max(0, start - retain + 1) // page_size evicted = [position for position in pages if position < first_retained_page] if evicted: - self.mem_manager.swa_page_allocator.free(torch.tensor([pages.pop(p) for p in evicted], dtype=torch.int32)) + self.mem_manager.swa_page_allocator.free([pages.pop(p) for p in evicted]) self.req_to_swa_pages[req_idx, :first_retained_page] = -1 missing = [p for p in range(start // page_size, (end + page_size - 1) // page_size) if p not in pages] if missing: @@ -140,7 +140,7 @@ def restore_state(self, req, state_cache_manager, buffer_idx, checkpoint_len): def clear_runtime_state(self, req_idx): pages = self._swa_pages[req_idx] if pages: - self.mem_manager.swa_page_allocator.free(torch.tensor(list(pages.values()), dtype=torch.int32)) + self.mem_manager.swa_page_allocator.free(list(pages.values())) pages.clear() self.req_to_swa_pages[req_idx].fill_(-1) diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 033851078b..04867ba586 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -1,6 +1,7 @@ import copy import importlib.util import json +import math import os import time @@ -9,7 +10,11 @@ from lightllm.common.basemodel.batch_objs import ModelInput from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_CPU_CACHE_TOKEN_PAGE_SIZE +from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import ( + DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, + DSV4_PROMPT_CACHE_PAGE_SIZE, + DSV4_SWA_PAGE_SIZE, +) from lightllm.models.deepseek_v4.layer_weights.pre_and_post_layer_weight import ( DeepseekV4PreAndPostLayerWeight, ) @@ -100,6 +105,16 @@ def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() state_mtp_step = 0 if self.args.run_mode == "prefill" else self.args.mtp_step reservations = self._get_post_profile_memory_reservations() + retain = self.config["sliding_window"] + DSV4_PROMPT_CACHE_PAGE_SIZE + decode_width = max(state_mtp_step + 1, 2 * state_mtp_step) + # The retained window may straddle one extra physical page. + per_req_pages = math.ceil((retain + decode_width - 1) / DSV4_SWA_PAGE_SIZE) + 1 + window_pages = self.max_req_num * per_req_pages + if self.args.run_mode != "prefill" and self.args.mtp_mode == "dspark": + window_pages += self.max_req_num + swa_page_num = window_pages + if self.args.run_mode != "decode": + swa_page_num += math.ceil(self.batch_max_tokens / DSV4_SWA_PAGE_SIZE) self.mem_manager = DeepseekV4MemoryManager( self.max_total_token_num, dtype=self.data_type, @@ -110,6 +125,7 @@ def _init_mem_manager(self): indexer_head_dim=self.config["index_head_dim"], max_request_num=self.max_req_num, mtp_step=state_mtp_step, + swa_page_num=swa_page_num, cpu_cache_token_page_size=( DSV4_CPU_CACHE_TOKEN_PAGE_SIZE if self.args.cpu_cache_token_page_size is None From 8157442131939f7bf8471e7b3ca87fbae7ecd8dd Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Thu, 24 Sep 2026 16:54:47 +0800 Subject: [PATCH 202/214] feat: isolate DeepSeek-V4 model support Keep the selected DSV4 model stack in one squashed change, with excluded commits documented in revert.md. Preserve the MTP hidden input buffer owned by decode CUDA Graph capture while retaining eager-mode cleanup. --- docs/CN/source/tutorial/api_server_args.rst | 27 +- docs/EN/source/tutorial/api_server_args.rst | 28 +- ...8192,topk=6,world_size=8}_NVIDIA_H200.json | 16 - ...8192,topk=6,world_size=8}_NVIDIA_H200.json | 16 - lightllm/common/basemodel/basemodel.py | 21 +- .../layer_infer/cache_tensor_manager.py | 7 +- .../meta_weights/fused_moe/ep_balance.py | 14 - .../meta_weights/fused_moe/ep_redundancy.py | 195 + .../meta_weights/fused_moe/eplb_placement.py | 468 --- .../fused_moe/expert_parallel_state.py | 56 - .../fused_moe/fused_moe_weight.py | 114 +- .../meta_weights/fused_moe/impl/__init__.py | 38 +- .../meta_weights/fused_moe/impl/base_impl.py | 91 +- .../fused_moe/impl/deepgemm_impl.py | 227 +- .../fused_moe/impl/marlin_impl.py | 4 - .../fused_moe/impl/triton_ep_impl.py | 42 - .../fused_moe/impl/triton_impl.py | 99 +- .../fused_moe/grouped_fused_moe_ep.py | 156 +- .../triton_kernel/fused_moe/grouped_topk.py | 251 -- .../fused_moe/sm90_fp8_triton_ep_moe.py | 575 --- .../triton_kernel/fused_moe/topk_select.py | 9 + .../quantization/fp8act_quant_kernel.py | 70 +- .../redundancy_topk_ids_repair.py | 111 + lightllm/common/eplb_utils.py | 16 - .../deepseek4_mem_manager.py | 12 +- .../kv_cache_mem_manager/mem_manager.py | 27 +- lightllm/common/quantization/deepgemm.py | 8 - lightllm/distributed/communication_op.py | 168 +- .../layer_infer/transformer_layer_infer.py | 17 +- .../layer_infer/transformer_layer_infer.py | 72 +- lightllm/models/deepseek_v4/model.py | 85 - .../triton_kernel/csrc/moe_topk_eplb.cu | 219 - .../deepseek_v4/triton_kernel/moe_topk.py | 81 - lightllm/models/gemma4/tokenizer.py | 2 - .../layer_infer/transformer_layer_infer.py | 17 +- lightllm/server/api_cli.py | 27 +- lightllm/server/api_http.py | 13 +- lightllm/server/api_http_pd.py | 13 +- lightllm/server/api_openai.py | 245 +- lightllm/server/api_start.py | 29 +- lightllm/server/api_stream_obj.py | 36 +- lightllm/server/build_prompt.py | 22 +- lightllm/server/core/objs/sampling_params.py | 3 +- lightllm/server/core/objs/start_args_type.py | 6 +- lightllm/server/function_call_parser.py | 10 +- lightllm/server/httpserver/async_queue.py | 9 - lightllm/server/httpserver/manager.py | 40 +- lightllm/server/httpserver/pd_loop.py | 66 +- .../httpserver_for_pd_master/manager.py | 119 +- .../pd_selector/cache_aware.py | 3 + .../pd_selector/prompt_cache_tree.py | 79 +- lightllm/server/metrics/metrics.py | 22 - .../multi_level_kv_cache/disk_cache_worker.py | 10 +- .../server/multi_level_kv_cache/manager.py | 24 +- lightllm/server/pd_io_struct.py | 43 - .../model_infer/mode_backend/base_backend.py | 12 +- .../mode_backend/chunked_prefill/impl.py | 5 - .../mode_backend/dp_backend/impl.py | 5 - .../mode_backend/ep_balance_monitor.py | 351 -- .../model_infer/mode_backend/eplb_manager.py | 546 --- .../model_infer/mode_backend/eplb_transfer.py | 759 ---- .../mode_backend/pd/base_kv_move_manager.py | 6 - .../decode_kv_move_manager.py | 7 +- .../decode_node_impl/decode_trans_process.py | 52 +- .../prefill_kv_move_manager.py | 7 +- .../prefill_trans_process.py | 5 - .../mode_backend/pd/trans_process_obj.py | 19 - .../mode_backend/redundancy_expert_manager.py | 158 + .../server/router/model_infer/model_rpc.py | 18 +- .../visualserver/model_infer/model_rpc.py | 3 +- lightllm/utils/device_utils.py | 5 - lightllm/utils/envs_utils.py | 73 +- lightllm/utils/error_utils.py | 4 - lightllm/utils/profile_max_tokens.py | 8 - lightllm/utils/start_utils.py | 111 +- revert.md | 47 + .../test_redundancy_expert_config.json | 180 + test/test_api/test_server_busy_handling.py | 19 +- .../test_pd_master_multi_choice.py | 6 +- .../basemodel/test_cuda_graph_layout.py | 42 + .../test_redundancy_topk_ids_repair.py | 151 + unit_tests/common/fused_moe/test_eplb.py | 3665 ----------------- .../fused_moe/test_eplb_transfer_gpu.py | 444 -- .../models/deepseek_v4/test_memory_profile.py | 114 - unit_tests/models/deepseek_v4/test_model.py | 57 + .../models/deepseek_v4/test_moe_topk.py | 91 - .../deepseek_v4/test_vision_integration.py | 22 +- unit_tests/models/gemma4/test_tokenizer.py | 38 - .../httpserver/test_pd_compact_transport.py | 1 - .../httpserver/test_pd_generate_error.py | 70 +- .../test_pd_master_cached_tokens.py | 2 +- .../test_running_request_lifecycle.py | 2 +- .../model_infer/test_ep_balance_monitor.py | 546 --- unit_tests/server/test_api_start_eplb.py | 64 - unit_tests/server/test_pd_master_mode.py | 17 - unit_tests/utils/test_envs_utils.py | 16 +- unit_tests/utils/test_start_utils.py | 5 +- 97 files changed, 1779 insertions(+), 10152 deletions(-) delete mode 100644 lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json delete mode 100644 lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py create mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py delete mode 100644 lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py create mode 100644 lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py delete mode 100644 lightllm/common/eplb_utils.py delete mode 100644 lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu delete mode 100644 lightllm/models/deepseek_v4/triton_kernel/moe_topk.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_manager.py delete mode 100644 lightllm/server/router/model_infer/mode_backend/eplb_transfer.py create mode 100644 lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py create mode 100644 revert.md create mode 100644 test/advanced_config/redundancy_expert/test_redundancy_expert_config.json create mode 100644 unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py delete mode 100644 unit_tests/common/fused_moe/test_eplb.py delete mode 100644 unit_tests/common/fused_moe/test_eplb_transfer_gpu.py delete mode 100644 unit_tests/models/deepseek_v4/test_memory_profile.py create mode 100644 unit_tests/models/deepseek_v4/test_model.py delete mode 100644 unit_tests/models/deepseek_v4/test_moe_topk.py delete mode 100644 unit_tests/models/gemma4/test_tokenizer.py delete mode 100644 unit_tests/server/router/model_infer/test_ep_balance_monitor.py delete mode 100644 unit_tests/server/test_api_start_eplb.py diff --git a/docs/CN/source/tutorial/api_server_args.rst b/docs/CN/source/tutorial/api_server_args.rst index f70d2a6c7b..1015e7b1d1 100644 --- a/docs/CN/source/tutorial/api_server_args.rst +++ b/docs/CN/source/tutorial/api_server_args.rst @@ -139,7 +139,7 @@ PD 分离模式参数 ``pd_node_resource_wait_timeout_seconds`` 为所有请求下发统一的资源等待上限;P/D 节点只负责按下发值 控制本地 ``shm_req`` 申请和 Router 等待进入推理系统,不读取本地限流开关或超时配置。首段的等待上限由 PD Master 上的 - ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 20 秒;设置为 -1 表示永久等待。 + ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 10 秒;设置为 -1 表示永久等待。 ``segment_index > 0`` 的续跑分段使用独立的等待上限,该值由 ``LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS`` 控制,默认 60 秒,以提高已产生部分结果的 请求最终完成的成功率。 @@ -150,15 +150,15 @@ PD 分离模式参数 则不再从头重试,以免产生重复内容。设置 ``--disable_pd_node_self_request_limit`` 后,PD Master 不再下发 有限的资源等待时间;P/D 节点永久等待,其他原因产生的 ``Server is busy`` 也会直接返回,不触发重试。 多机 TP 场景仅由 master 节点执行超时判断,slave 节点永久等待。cache 命中记录允许提升优先级的最大年龄由 - ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` 控制,默认 180 秒。cache 命中提权还要求输入 - token 数达到 ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` 配置的门槛(默认 2048),避免短请求仅因 + ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` 控制,默认 36 秒。cache 命中提权还要求输入 + token 数达到 ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` 配置的门槛(默认 4096),避免短请求仅因 cache 命中率高而提升优先级。 启动示例: .. code-block:: bash - LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=20 \ + LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=10 \ LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS=60 \ LIGHTLLM_PD_NODE_BUSY_RETRY_TIMEOUT_SECONDS=120 \ python -m lightllm.server.api_server --run_mode pd_master ... @@ -238,10 +238,10 @@ PD 分离模式参数 .. option:: --running_max_req_size - 本机共享请求槽总数,默认为 ``256``。DP 模式下,每个 rank 的调度并发和 - CUDA Graph batch 上限按该值除以本机 DP rank 数并向下取整。 - 开启 DP prompt cache fetch 时,普通或 Prefill 节点的模型请求槽仍保留全局容量, - 供跨 rank 前缀匹配使用。PD Master 当前不使用该参数执行请求准入限流。 + 本机共享请求槽总数,默认为 ``256``。 + DP 模式按 ``ceil(running_max_req_size / 本机 DP rank 数)`` + 分配每个 rank 的请求状态槽,调度并发和 CUDA Graph batch 上限也受此限制。 + diverse 模式及 DP prompt cache fetch 模式保持原有容量。 .. option:: --max_req_total_len @@ -735,6 +735,17 @@ MTP 多预测参数 增加此值允许更多预测,但确保模型与指定的步数兼容。 目前 deepseekv3/r1 模型仅支持 1 步 +DeepSeek 冗余专家参数 +--------------------- + +.. option:: --ep_redundancy_expert_config_path + + 冗余专家配置的路径。可用于 deepseekv3 模型。 + +.. option:: --auto_update_redundancy_expert + + 是否通过在线专家使用计数器为 deepseekv3 模型更新冗余专家。 + 监控和日志参数 -------------- diff --git a/docs/EN/source/tutorial/api_server_args.rst b/docs/EN/source/tutorial/api_server_args.rst index 35465051b6..a58c607971 100644 --- a/docs/EN/source/tutorial/api_server_args.rst +++ b/docs/EN/source/tutorial/api_server_args.rst @@ -147,7 +147,7 @@ PD disaggregation Mode Parameters ``pd_node_resource_wait_timeout_seconds`` for every request. P/D nodes only enforce the received value for local ``shm_req`` allocation and the wait from Router entry to inference entry; they do not read local limiting switches or timeout settings. The first segment's timeout is - controlled on PD Master by ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 20 seconds; set it + controlled on PD Master by ``LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 10 seconds; set it to -1 to wait indefinitely. Continuation segments with ``segment_index > 0`` use a separate timeout controlled by ``LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS`` and defaults to 60 seconds, improving the chance that requests which have already produced partial results complete successfully. When set to a non-negative value, @@ -162,15 +162,15 @@ PD disaggregation Mode Parameters In multi-node TP deployments, only the master node evaluates the timeout; slave nodes wait indefinitely. The maximum cache-record age eligible for promotion is controlled by ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS`` and defaults to - 180 seconds. Cache-hit promotion also requires at least the number of input tokens configured by - ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` (2048 by default), so short requests do not gain priority + 36 seconds. Cache-hit promotion also requires at least the number of input tokens configured by + ``LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS`` (4096 by default), so short requests do not gain priority solely from a high cache-hit rate. Startup example: .. code-block:: bash - LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=20 \ + LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS=10 \ LIGHTLLM_PD_NODE_CONTINUATION_RESOURCE_WAIT_TIMEOUT_SECONDS=60 \ LIGHTLLM_PD_NODE_BUSY_RETRY_TIMEOUT_SECONDS=120 \ python -m lightllm.server.api_server --run_mode pd_master ... @@ -253,11 +253,10 @@ Memory and Batch Processing Parameters .. option:: --running_max_req_size - Total shared request slots on the local node, default ``256``. In DP mode, - each rank's scheduling concurrency and CUDA Graph batch limit use this value - divided by the local DP rank count, rounded down. With DP prompt cache fetch, - Normal and Prefill model request slots retain the global capacity for cross-rank - prefix matching. PD Master does not currently use this parameter for admission limiting. + Total shared request slots on the local node, default is ``256``. + DP mode allocates ``ceil(running_max_req_size / local DP rank count)`` + request-state slots per rank, which also limits scheduling concurrency and CUDA Graph batches. + Diverse mode and DP prompt cache fetch keep the original capacity. .. option:: --max_req_total_len @@ -752,6 +751,17 @@ MTP Multi-Prediction Parameters Increasing this value allows more predictions, but ensure the model is compatible with the specified number of steps. Currently deepseekv3/r1 models only support 1 step +DeepSeek Redundant Expert Parameters +------------------------------------ + +.. option:: --ep_redundancy_expert_config_path + + Path to redundant expert configuration. Can be used for deepseekv3 models. + +.. option:: --auto_update_redundancy_expert + + Whether to update redundant experts for deepseekv3 models through online expert usage counters. + Monitoring and Logging Parameters --------------------------------- diff --git a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json deleted file mode 100644 index 694771f831..0000000000 --- a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=128,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json +++ /dev/null @@ -1,16 +0,0 @@ -{ - "publish_rows_large": 32, - "publish_rows_small": 4, - "publish_num_warps": 2, - "histogram_block": 2048, - "histogram_num_warps": 8, - "pull_num_warps": 8, - "activation_programs": 1024, - "activation_rows": 32, - "activation_num_warps": 8, - "return_num_warps": 8, - "combine_rows_large": 2, - "combine_rows_small": 1, - "combine_block": 512, - "combine_num_warps": 4 -} diff --git a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json b/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json deleted file mode 100644 index 625756625c..0000000000 --- a/lightllm/common/all_kernel_configs/sm90_fp8_triton_ep_moe/{alignment=256,hidden_size=4096,intermediate_size=2048,local_experts=32,num_max_tokens_per_rank=8192,topk=6,world_size=8}_NVIDIA_H200.json +++ /dev/null @@ -1,16 +0,0 @@ -{ - "publish_rows_large": 32, - "publish_rows_small": 4, - "publish_num_warps": 8, - "histogram_block": 512, - "histogram_num_warps": 2, - "pull_num_warps": 8, - "activation_programs": 1024, - "activation_rows": 32, - "activation_num_warps": 8, - "return_num_warps": 8, - "combine_rows_large": 2, - "combine_rows_small": 1, - "combine_block": 512, - "combine_num_warps": 4 -} diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 6358e4327e..6e8e54dfad 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -73,8 +73,6 @@ class TpPartBaseModel: def __init__(self, kvargs): self.args = get_env_start_args() - self.eplb_manager = None - self.ep_balance_monitor = None self.run_mode = kvargs["run_mode"] self.weight_dir_ = kvargs["weight_dir"] self.max_total_token_num = kvargs["max_total_token_num"] @@ -333,16 +331,9 @@ def forward(self, model_input: ModelInput): model_input.to_cuda() if model_input.is_prefill: - model_output = self._prefill(model_input) - self._after_prefill() - return model_output - return self._decode(model_input) - - def _after_prefill(self): - if self.ep_balance_monitor is not None: - self.ep_balance_monitor.record_prefill_round() - if self.eplb_manager is not None: - self.eplb_manager.step() + return self._prefill(model_input) + else: + return self._decode(model_input) def _select_mem_indexes(self, model_input: ModelInput): if model_input.is_prefill: @@ -717,7 +708,8 @@ def _token_forward(self, infer_state: InferStateInfo): 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) - infer_state.mtp_draft_input_hiddens = None + if not infer_state.is_cuda_graph: + infer_state.mtp_draft_input_hiddens = None input_embs = self.pre_infer._tpsp_sp_split(input=input_embs, infer_state=infer_state) for i in range(self.layers_num): @@ -806,7 +798,6 @@ def _microbatch_overlap_prefill_cuda(self, model_input0: ModelInput, model_input dist_group_manager.clear_deepep_buffer() model_output0.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event model_output1.prefill_mem_indexes_ready_event = prefill_mem_indexes_ready_event - self._after_prefill() return model_output0, model_output1 @torch.no_grad() @@ -1067,7 +1058,7 @@ def _autotune_warmup(self): warmup_max_tokens = self.batch_max_tokens if self.args.run_mode == "decode": - decode_rows = self.max_req_num * self.mtp_manager.get_decode_batch_multiplier(self.is_mtp_draft_model) + decode_rows = self.max_req_num * self.mtp_manager.get_decode_tokens_per_request(self.is_mtp_draft_model) warmup_max_tokens = min(warmup_max_tokens, decode_rows) warmup_lengths = [1, 4, 8, 16, 32, 64, 128, 256, 1024, 2048, 4096] diff --git a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py index 13906d0d8a..8bcf99b992 100644 --- a/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py +++ b/lightllm/common/basemodel/layer_infer/cache_tensor_manager.py @@ -22,12 +22,7 @@ def custom_del(self: torch.Tensor): if hasattr(self, "storage_weak_ptr"): storage_weak_ptr = self.storage_weak_ptr else: - try: - storage_weak_ptr = self.untyped_storage()._weak_ref() - except RuntimeError: - # Some tensor implementations, including UndefinedTensorImpl, - # have no backing storage. Their destructor must stay silent. - return + storage_weak_ptr = self.untyped_storage()._weak_ref() UntypedStorage._free_weak_ref(storage_weak_ptr) if storage_weak_ptr in g_cache_manager.ptr_to_bufnode: g_cache_manager.changed_ptr.add(storage_weak_ptr) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py deleted file mode 100644 index 70435146c5..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_balance.py +++ /dev/null @@ -1,14 +0,0 @@ -from dataclasses import dataclass - - -@dataclass(slots=True) -class PrefillEPBalanceCounters: - """Cumulative CPU loads for one EP MoE layer's completed prefill dispatches.""" - - route_load: int = 0 - compute_load: int = 0 - - def accumulate(self, route_load: int, compute_load: int): - """Accumulate exact route and alignment-expanded compute loads for one prefill dispatch.""" - self.route_load += route_load - self.compute_load += compute_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py new file mode 100644 index 0000000000..749400c8d8 --- /dev/null +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/ep_redundancy.py @@ -0,0 +1,195 @@ +import numpy as np +import torch +from .fused_moe_weight import FusedMoeWeight +from lightllm.utils.log_utils import init_logger +from typing import Dict + +logger = init_logger(__name__) + + +class FusedMoeWeightEPAutoRedundancy: + def __init__( + self, + ep_fused_moe_weight: FusedMoeWeight, + ) -> None: + super().__init__() + self._ep_w = ep_fused_moe_weight + self.redundancy_expert_num = self._ep_w.redundancy_expert_num + + def clear_counter(self): + self._ep_w.routed_expert_counter_tensor.fill_(0) + return + + def prepare_redundancy_experts( + self, + ): + expert_counter = self._ep_w.routed_expert_counter_tensor.detach().cpu().numpy() + logger.info( + f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" + f" expert_counter: {expert_counter}" + ) + self._ep_w.routed_expert_counter_tensor.fill_(0) + ep_n_routed_experts = self._ep_w.n_routed_experts // self._ep_w.global_world_size + start_expert_id = ep_n_routed_experts * self._ep_w.global_rank_ + no_redundancy_expert_ids = list(range(start_expert_id, start_expert_id + ep_n_routed_experts)) + + # 统计 0 rank 上的全局 topk 冗余信息,帮助导出一份全局可用的静态使用的冗余专家静态配置。 + if self._ep_w.global_rank_ == 0: + # int(e) for serialization, int64 can not be serialized by json.dump. + topk_redundancy_expert_ids = list(int(e) for e in np.argsort(expert_counter)[-self.redundancy_expert_num :]) + else: + topk_redundancy_expert_ids = None + + # 不要选中当前已经存在的非冗余专家作为冗余专家 + expert_counter[no_redundancy_expert_ids] = 0 + + self.redundancy_expert_ids = list(np.argsort(expert_counter)[-self.redundancy_expert_num :]) + logger.info( + f"layer_index {self._ep_w.layer_num_} global_rank {self._ep_w.global_rank_}" + f" new select redundancy_expert_ids : {self.redundancy_expert_ids}" + ) + + # 准备加载过度变量。 + self.experts_up_projs = [None] * self.redundancy_expert_num + self.experts_gate_projs = [None] * self.redundancy_expert_num + self.experts_up_proj_scales = [None] * self.redundancy_expert_num + self.experts_gate_proj_scales = [None] * self.redundancy_expert_num + self.w2_list = [None] * self.redundancy_expert_num + self.w2_scale_list = [None] * self.redundancy_expert_num + self.w13 = [None, None] # weight, weight_scale + self.w2 = [None, None] # weight, weight_scale + return topk_redundancy_expert_ids + + def load_hf_weights(self, weights): + # 加载冗余专家的权重参数 + for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): + i_experts = redundant_expert_id + w1_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.weight" + w2_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.weight" + w3_weight = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.weight" + if w1_weight in weights: + self.experts_gate_projs[i] = weights[w1_weight] + if w3_weight in weights: + self.experts_up_projs[i] = weights[w3_weight] + if w2_weight in weights: + self.w2_list[i] = weights[w2_weight] + + self._load_weight_scale(weights) + self._fuse() + + def _fuse(self): + self._fuse_weight_scale() + + with self._ep_w.lock: + if ( + hasattr(self, "experts_up_projs") + and None not in self.experts_up_projs + and None not in self.experts_gate_projs + and None not in self.w2_list + ): + gate_out_dim, gate_in_dim = self.experts_gate_projs[0].shape + up_out_dim, up_in_dim = self.experts_up_projs[0].shape + assert gate_in_dim == up_in_dim + dtype = self.experts_gate_projs[0].dtype + total_expert_num = self.redundancy_expert_num + + w13 = torch.empty((total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu") + + for i_experts in range(self.redundancy_expert_num): + w13[i_experts, 0:gate_out_dim:, :] = self.experts_gate_projs[i_experts] + w13[i_experts, gate_out_dim:, :] = self.experts_up_projs[i_experts] + + inter_shape, hidden_size = self.w2_list[0].shape[0], self.w2_list[0].shape[1] + w2 = torch._utils._flatten_dense_tensors(self.w2_list).view(len(self.w2_list), inter_shape, hidden_size) + if self._ep_w.quant_method._check_weight_need_quanted(weight=w13): + w13_pack, _ = self._ep_w.quant_method.create_moe_weight( + out_dims=[gate_out_dim + up_out_dim], + in_dim=1, + dtype=self._ep_w.data_type_, + device_id=self._ep_w.device_id_, + num_experts=self.redundancy_expert_num, + ) + self._ep_w.quant_method.quantize(w13, w13_pack) + w2_pack, _ = self._ep_w.quant_method.create_moe_weight( + out_dims=[inter_shape], + in_dim=hidden_size, + dtype=self._ep_w.data_type_, + device_id=self._ep_w.device_id_, + num_experts=self.redundancy_expert_num, + ) + self._ep_w.quant_method.quantize(w2, w2_pack) + + self.w13[0] = w13_pack.weight + self.w13[1] = w13_pack.weight_scale + self.w2[0] = w2_pack.weight + self.w2[1] = w2_pack.weight_scale + else: + self.w13[0] = w13 + self.w2[0] = w2 + delattr(self, "w2_list") + delattr(self, "experts_up_projs") + delattr(self, "experts_gate_projs") + + def _fuse_weight_scale(self): + with self._ep_w.lock: + if ( + hasattr(self, "experts_up_proj_scales") + and None not in self.experts_up_proj_scales + and None not in self.experts_gate_proj_scales + and None not in self.w2_scale_list + ): + gate_out_dim, gate_in_dim = self.experts_gate_proj_scales[0].shape + up_out_dim, up_in_dim = self.experts_up_proj_scales[0].shape + assert gate_in_dim == up_in_dim + dtype = self.experts_gate_proj_scales[0].dtype + total_expert_num = self.redundancy_expert_num + w13_scale = torch.empty( + (total_expert_num, gate_out_dim + up_out_dim, gate_in_dim), dtype=dtype, device="cpu" + ) + for i_experts in range(self.redundancy_expert_num): + w13_scale[i_experts, 0:gate_out_dim:, :] = self.experts_gate_proj_scales[i_experts] + w13_scale[i_experts, gate_out_dim:, :] = self.experts_up_proj_scales[i_experts] + + inter_shape, hidden_size = self.w2_scale_list[0].shape[0], self.w2_scale_list[0].shape[1] + w2_scale = torch._utils._flatten_dense_tensors(self.w2_scale_list).view( + len(self.w2_scale_list), inter_shape, hidden_size + ) + self.w13[1] = w13_scale + self.w2[1] = w2_scale + delattr(self, "w2_scale_list") + delattr(self, "experts_up_proj_scales") + delattr(self, "experts_gate_proj_scales") + + def _load_weight_scale(self, weights: Dict[str, torch.Tensor]) -> None: + # 加载冗余专家的scale参数 + for i, redundant_expert_id in enumerate(self.redundancy_expert_ids): + i_experts = redundant_expert_id + weight_scale_suffix = self._ep_w.quant_method.weight_scale_suffix + w1_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w1_weight_name}.{weight_scale_suffix}" + w2_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w2_weight_name}.{weight_scale_suffix}" + w3_scale = f"{self._ep_w.weight_prefix}.{i_experts}.{self._ep_w.w3_weight_name}.{weight_scale_suffix}" + if w1_scale in weights: + self.experts_gate_proj_scales[i] = weights[w1_scale] + if w3_scale in weights: + self.experts_up_proj_scales[i] = weights[w3_scale] + if w2_scale in weights: + self.w2_scale_list[i] = weights[w2_scale] + + def commit(self): + for index, dest_tensor in enumerate([self._ep_w.w13.weight, self._ep_w.w13.weight_scale]): + if dest_tensor is not None: + assert isinstance( + dest_tensor, torch.Tensor + ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" + dest_tensor[-self.redundancy_expert_num :, :, :] = self.w13[index][:, :, :] + + for index, dest_tensor in enumerate([self._ep_w.w2.weight, self._ep_w.w2.weight_scale]): + if dest_tensor is not None: + assert isinstance( + dest_tensor, torch.Tensor + ), f"dest_tensor should be a torch.Tensor, but got {type(dest_tensor)}" + dest_tensor[-self.redundancy_expert_num :, :, :] = self.w2[index][:, :, :] + + self._ep_w.redundancy_expert_ids_tensor.copy_( + torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cpu") + ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py deleted file mode 100644 index 1ea8eb1643..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/eplb_placement.py +++ /dev/null @@ -1,468 +0,0 @@ -from dataclasses import dataclass -from functools import lru_cache -from typing import Dict, Tuple -import torch - - -def build_initial_redundant_expert_ids( - num_logical_experts: int, - num_ranks: int, - num_redundant_experts_per_rank: int, -) -> torch.Tensor: - """Build a deterministic initial placement without local duplicates.""" - assert num_logical_experts % num_ranks == 0 - num_experts_per_rank = num_logical_experts // num_ranks - assert 0 < num_redundant_experts_per_rank <= num_logical_experts - num_experts_per_rank - - # 初始化结果确定,不依赖随机数。 - # 每个 rank 不会复制自己原本拥有的 expert。 - # 同一个 rank 的冗余槽位不会重复。 - # 最后一个 rank 通过取模自然回绕。 - rank_offsets = torch.arange(1, num_ranks + 1, dtype=torch.int64)[:, None] * num_experts_per_rank - expert_offsets = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64) - return (rank_offsets + expert_offsets) % num_logical_experts - - -def build_logical_to_physical_map( - redundant_expert_ids: torch.Tensor, # 冗余布局,shape 为 [num_ranks, num_redundant_experts_per_rank]。 - num_logical_experts: int, # 逻辑 expert 的总数。 - source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 - node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 -) -> Tuple[ - torch.Tensor, torch.Tensor -]: # logical_to_physical [num_logical_experts, num_ranks], replica_counts [num_logical_experts] - """构建单层逻辑 expert 到物理副本的映射""" - - logical_to_physical, replica_counts = build_logical_to_physical_maps_for_layers( - redundant_expert_ids.unsqueeze(0), - num_logical_experts, - source_rank=source_rank, - node_world_size=node_world_size, - ) - return logical_to_physical.squeeze(0), replica_counts.squeeze(0) - - -def build_logical_to_physical_maps_for_layers( - redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] - num_logical_experts: int, # 逻辑 expert 的总数。 - source_rank: int | None = None, # 可选的全局源 rank;传入时优先选择同节点副本。 - node_world_size: int | None = None, # 单个节点包含的 rank 数;source_rank 非空时必填。 -) -> Tuple[ - torch.Tensor, # logical_to_physical, shape [num_layers, num_logical_experts, num_ranks] - torch.Tensor, # replica_counts, shape [num_layers, num_logical_experts] -]: - """Build stable CPU int32 maps for the supplied layers without modifying the input.""" - if redundant_expert_ids_by_layer.ndim != 3: - raise ValueError("redundant_expert_ids_by_layer must be [layers, ranks, num_redundant_experts_per_rank]") - num_ranks, num_redundant_experts_per_rank = redundant_expert_ids_by_layer.shape[1:] - assert num_logical_experts % num_ranks == 0 - layout = _get_physical_expert_layout(num_logical_experts, num_ranks, num_redundant_experts_per_rank) - logical_to_physical, replica_counts = _build_global_replica_maps_for_layers(redundant_expert_ids_by_layer, layout) - if source_rank is None: - return logical_to_physical, replica_counts - - assert node_world_size is not None - replica_positions = torch.arange(num_ranks, dtype=torch.int64) - compact_maps_by_layer, selected_counts_by_layer = _select_source_node_replicas( - logical_to_physical, - replica_counts, - source_rank=source_rank, - node_world_size=node_world_size, - num_physical_experts_per_rank=layout.num_physical_experts_per_rank, - replica_positions=replica_positions, - ) - return ( - _rotate_selected_replicas( - compact_maps_by_layer, - selected_counts_by_layer, - source_rank=source_rank, - replica_positions=replica_positions, - ), - selected_counts_by_layer, - ) - - -def select_improving_placements( - expert_load: torch.Tensor, - current_placement: torch.Tensor, - candidate_placement: torch.Tensor, - *, - rebalance_gain_threshold: float, - expert_alignment: int | None = None, - node_world_size: int | None = None, -) -> Tuple[torch.Tensor, torch.Tensor, Dict[str, float | int], torch.Tensor, torch.Tensor]: - """Select better layers and return current/final rank loads without re-estimation.""" - if not 0.0 <= rebalance_gain_threshold <= 1.0: - raise ValueError("rebalance_gain_threshold must be between 0.0 and 1.0") - assert current_placement.shape == candidate_placement.shape - current_rank_load = _estimate_rank_load(expert_load, current_placement, expert_alignment, node_world_size) - candidate_rank_load = _estimate_rank_load(expert_load, candidate_placement, expert_alignment, node_world_size) - current_critical = current_rank_load.max(dim=-1).values.sum(dim=0) - candidate_critical = candidate_rank_load.max(dim=-1).values.sum(dim=0) - # Each changed layer must reduce its own critical load. All selected - # changes must then collectively meet the configured model-level - # critical-load reduction threshold, avoiding low-gain migrations. - improved = candidate_critical < current_critical - selected = current_placement.clone() - selected[improved] = candidate_placement[improved] - selected_rank_load = torch.where(improved[:, None], candidate_rank_load, current_rank_load) - model_current_critical = current_critical.sum() - model_current_mean = current_rank_load.mean(dim=-1).sum() - model_selected_critical = selected_rank_load.max(dim=-1).values.sum() - model_selected_mean = selected_rank_load.mean(dim=-1).sum() - model_ratio = model_current_critical / model_current_mean.clamp_min(1.0) - candidate_model_ratio = model_selected_critical / model_selected_mean.clamp_min(1.0) - candidate_rebalance_gain = (model_current_critical - model_selected_critical) / model_current_critical.clamp_min( - 1.0 - ) - metrics = { - "model_imbalance_ratio": float(model_ratio.item()), - "candidate_model_imbalance_ratio": float(candidate_model_ratio.item()), - "candidate_rebalance_gain": float(candidate_rebalance_gain.item()), - "candidate_changed_layer_count": int(improved.sum().item()), - } - if candidate_rebalance_gain >= rebalance_gain_threshold: - return selected, improved, metrics, current_rank_load, selected_rank_load - return ( - current_placement.clone(), - torch.zeros_like(improved), - metrics, - current_rank_load, - current_rank_load, - ) - - -def plan_redundant_experts( - expert_load: torch.Tensor, - num_ranks: int, - num_redundant_experts_per_rank: int, - expert_alignment: int | None = None, - node_world_size: int | None = None, - current_placement: torch.Tensor | None = None, - stickiness: float = 0.0, -) -> torch.Tensor: - """Plan replicas from [samples, layers, source_nodes, experts] loads. - - With ``current_placement`` and positive ``stickiness``, a candidate that - keeps an expert on its current rank receives a bonus of - ``stickiness * mean per-layer expert load``. This preserves rank - membership, not a particular redundant physical slot; target slots are - canonicalized against the current live rows before transfer and metadata - publication. A rank membership only changes when the move improves the - critical-load objective by more than that margin. - With zero stickiness, placement is determined solely by the load objective. - """ - if expert_alignment is not None: - assert expert_alignment > 0 - node_world_size = _resolve_node_world_size(expert_load, num_ranks, node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape - assert num_logical_experts % num_ranks == 0 - assert num_redundant_experts_per_rank > 0 - num_experts_per_rank = num_logical_experts // num_ranks - num_redundant = num_ranks * num_redundant_experts_per_rank - assert num_redundant <= num_logical_experts * (num_ranks - 1) - - load = expert_load.to(dtype=torch.float64, device="cpu") - placement = torch.full((num_layers, num_ranks, num_redundant_experts_per_rank), -1, dtype=torch.int64) - owner_rank = torch.arange(num_logical_experts, dtype=torch.int64) // num_experts_per_rank - if current_placement is not None: - assert tuple(current_placement.shape) == ( - num_layers, - num_ranks, - num_redundant_experts_per_rank, - ) - current_locations = _expert_locations(current_placement, num_logical_experts) - stickiness_scale = load.sum(dim=(0, 2, 3)) / num_logical_experts - else: - current_locations = None - stickiness_scale = None - - locations = _expert_locations(placement, num_logical_experts) - expert_rank = _expert_rank_load_all(load, locations, num_nodes, node_world_size, expert_alignment) - rank_load = expert_rank.sum(dim=2) - remaining_slots = torch.full((num_layers, num_ranks), num_redundant_experts_per_rank, dtype=torch.int64) - layer_indices = torch.arange(num_layers, dtype=torch.int64) - expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - # Every iteration fills one slot per layer. Candidate expert evaluation - # is vectorized across all layers and logical experts, which keeps large - # GLM/Qwen planning comfortably on the CPU fast path. - for _ in range(num_redundant): - rank_order = torch.argsort(rank_load.sum(dim=0), dim=1, stable=True) - target_ranks = torch.full((num_layers,), -1, dtype=torch.int64) - legal = torch.zeros((num_layers, num_logical_experts), dtype=torch.bool) - for layer in range(num_layers): - for target_rank in rank_order[layer].tolist(): - if remaining_slots[layer, target_rank] == 0: - continue - candidate_legal = (owner_rank != target_rank) & ~locations[layer, :, target_rank] - if torch.any(candidate_legal): - target_ranks[layer] = target_rank - legal[layer] = candidate_legal - break - if torch.any(target_ranks < 0): - raise RuntimeError("EPLB planner found no valid redundant expert placement") - - candidate_locations = locations.clone() - candidate_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] = True - candidate_expert_rank = _expert_rank_load_all( - load, candidate_locations, num_nodes, node_world_size, expert_alignment - ) - candidate_rank_load = rank_load[:, :, None, :] - expert_rank + candidate_expert_rank - critical = candidate_rank_load.max(dim=3).values.sum(dim=0) - critical.masked_fill_(~legal, torch.inf) - if current_locations is not None: - # An expert already held by the target rank is retained unless - # another candidate beats it by more than the stickiness margin. - # This is rank membership, not physical-slot stickiness. Masked - # (inf) candidates stay masked: inf - x == inf. - keep = current_locations[layer_indices[:, None], expert_ids[None, :], target_ranks[:, None]] - critical = critical - stickiness * stickiness_scale[:, None] * keep - selected_experts = critical.argmin(dim=1) - if torch.isinf(critical[layer_indices, selected_experts]).any(): - raise RuntimeError("EPLB planner found no valid redundant expert placement") - - slots = num_redundant_experts_per_rank - remaining_slots[layer_indices, target_ranks] - placement[layer_indices, target_ranks, slots] = selected_experts - selected_next = candidate_expert_rank[:, layer_indices, selected_experts] - selected_old = expert_rank[:, layer_indices, selected_experts] - rank_load += selected_next - selected_old - expert_rank[:, layer_indices, selected_experts] = selected_next - locations[layer_indices, selected_experts, target_ranks] = True - remaining_slots[layer_indices, target_ranks] -= 1 - - assert torch.all(placement >= 0) - return placement - - -@dataclass(frozen=True, eq=False) -class _PhysicalExpertLayout: - """进程内按拓扑复用的只读物理 expert 布局;其中 Tensor 不得原地修改。""" - - num_logical_experts: int - num_ranks: int - num_physical_experts_per_rank: int - primary_physical_ids: torch.Tensor - redundant_physical_ids: torch.Tensor - - -@lru_cache(maxsize=8) -def _get_physical_expert_layout( - num_logical_experts: int, - num_ranks: int, - num_redundant_experts_per_rank: int, -) -> _PhysicalExpertLayout: - """返回按静态拓扑缓存的只读 CPU 物理 expert ID。""" - num_experts_per_rank = num_logical_experts // num_ranks - num_physical_experts_per_rank = num_experts_per_rank + num_redundant_experts_per_rank - expert_ids = torch.arange(num_logical_experts, dtype=torch.int64) - primary_physical_ids = ( - (expert_ids // num_experts_per_rank) * num_physical_experts_per_rank + expert_ids % num_experts_per_rank - ).to(torch.int32) - ranks = torch.arange(num_ranks, dtype=torch.int64).repeat_interleave(num_redundant_experts_per_rank) - slots = torch.arange(num_redundant_experts_per_rank, dtype=torch.int64).repeat(num_ranks) - redundant_physical_ids = (ranks * num_physical_experts_per_rank + num_experts_per_rank + slots).to(torch.int32) - return _PhysicalExpertLayout( - num_logical_experts=num_logical_experts, - num_ranks=num_ranks, - num_physical_experts_per_rank=num_physical_experts_per_rank, - primary_physical_ids=primary_physical_ids, - redundant_physical_ids=redundant_physical_ids, - ) - - -def _build_global_replica_maps_for_layers( - redundant_expert_ids_by_layer: torch.Tensor, # [num_layers, num_ranks, num_redundant_experts_per_rank] - layout: _PhysicalExpertLayout, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Build stable global maps with primary copies first and unused slots set to ``-1``.""" - - num_layers = redundant_expert_ids_by_layer.shape[0] - num_logical_experts = layout.num_logical_experts - max_replicas = layout.num_ranks - redundant_ids = redundant_expert_ids_by_layer.to(dtype=torch.int64, device="cpu") - logical_to_physical = torch.full((num_layers, num_logical_experts, max_replicas), -1, dtype=torch.int32) - logical_to_physical[:, :, 0] = layout.primary_physical_ids - replica_counts = torch.ones((num_layers, num_logical_experts), dtype=torch.int32) - - flat_redundant_ids = redundant_ids.reshape(num_layers, -1) - if not flat_redundant_ids.numel(): - return logical_to_physical, replica_counts - - # 稳定排序保留 rank-major、slot-major 的历史顺序;第 0 列固定为主副本。 - sort_order = torch.argsort(flat_redundant_ids, dim=1, stable=True) - sorted_redundant_ids = flat_redundant_ids.gather(1, sort_order) - flat_positions = torch.arange(flat_redundant_ids.shape[1], dtype=torch.int64).unsqueeze(0) - group_starts = torch.where( - torch.cat( - ( - torch.ones((num_layers, 1), dtype=torch.bool), - sorted_redundant_ids[:, 1:] != sorted_redundant_ids[:, :-1], - ), - dim=1, - ), - flat_positions, - 0, - ) - replica_indices = flat_positions - torch.cummax(group_starts, dim=1).values + 1 - redundant_counts = torch.zeros((num_layers, num_logical_experts), dtype=torch.int32) - redundant_counts.scatter_add_( - 1, - flat_redundant_ids, - torch.ones_like(flat_redundant_ids, dtype=torch.int32), - ) - assert int(redundant_counts.max().item()) < max_replicas, "an expert can have at most one replica per rank" - replica_counts += redundant_counts - - layer_indices = torch.arange(num_layers, dtype=torch.int64).view(-1, 1).expand_as(sort_order) - redundant_physical_ids = layout.redundant_physical_ids.unsqueeze(0).expand_as(sort_order).gather(1, sort_order) - logical_to_physical[layer_indices, sorted_redundant_ids, replica_indices] = redundant_physical_ids - return logical_to_physical, replica_counts - - -def _select_source_node_replicas( - logical_to_physical: torch.Tensor, - replica_counts: torch.Tensor, - *, - source_rank: int, - node_world_size: int, - num_physical_experts_per_rank: int, - replica_positions: torch.Tensor, -) -> Tuple[torch.Tensor, torch.Tensor]: - """Keep source-node replicas when available, otherwise fall back to all stable candidates.""" - num_layers, num_logical_experts, _max_replicas = logical_to_physical.shape - source_node = source_rank // node_world_size - output_positions = replica_positions.view(1, 1, -1) - valid = output_positions < replica_counts.unsqueeze(-1) - local = valid & ( - torch.div( - logical_to_physical, - num_physical_experts_per_rank * node_world_size, - rounding_mode="floor", - ) - == source_node - ) - selected = torch.where(local.any(dim=2, keepdim=True), local, valid) - selected_counts_by_layer = selected.sum(dim=2, dtype=torch.int32) - - compact_maps_by_layer = torch.full_like(logical_to_physical, -1) - selected_positions = selected.cumsum(dim=2) - 1 - layers = torch.arange(num_layers, dtype=torch.int64).view(-1, 1, 1).expand_as(selected) - experts = torch.arange(num_logical_experts, dtype=torch.int64).view(1, -1, 1).expand_as(selected) - compact_maps_by_layer[layers[selected], experts[selected], selected_positions[selected]] = logical_to_physical[ - selected - ] - return compact_maps_by_layer, selected_counts_by_layer - - -def _rotate_selected_replicas( - compact_maps_by_layer: torch.Tensor, - selected_count_by_layer: torch.Tensor, - *, - source_rank: int, - replica_positions: torch.Tensor, -) -> torch.Tensor: - """根据 source_rank 循环调整每个逻辑 expert 的候选副本顺序,让不同源 rank 优先使用不同副本,同时保持候选副本集合和副本数量不变""" - output_positions = replica_positions.view(1, 1, -1) - selected_count64_by_layer = selected_count_by_layer.to(torch.int64).unsqueeze(-1) - source_positions_by_layer = (output_positions + source_rank) % selected_count64_by_layer - maps_by_layer = compact_maps_by_layer.gather(2, source_positions_by_layer) - maps_by_layer.masked_fill_(output_positions >= selected_count64_by_layer, -1) - return maps_by_layer - - -def _estimate_rank_load( - expert_load: torch.Tensor, - redundant_expert_ids: torch.Tensor, - expert_alignment: int | None = None, - node_world_size: int | None = None, -) -> torch.Tensor: - """Estimate [samples, layers, ranks] load from source-node-local routing. - - Source loads remain separate until assigned to physical replicas, then - combine before the per-expert alignment used by DeepEP. - """ - node_world_size = _resolve_node_world_size(expert_load, redundant_expert_ids.shape[1], node_world_size) - num_samples, num_layers, num_nodes, num_logical_experts = expert_load.shape - assert redundant_expert_ids.ndim == 3 and redundant_expert_ids.shape[0] == num_layers - num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape[1:] - assert num_logical_experts % num_ranks == 0 - if expert_alignment is not None: - assert expert_alignment > 0 - - rank_load = _expert_rank_load_all( - expert_load, - _expert_locations(redundant_expert_ids, num_logical_experts), - num_nodes, - node_world_size, - expert_alignment, - ).sum(dim=2) - return rank_load - - -def _resolve_node_world_size(expert_load: torch.Tensor, num_ranks: int, node_world_size: int | None) -> int: - """Validate production [samples, layers, source_nodes, experts] planner loads.""" - assert expert_load.ndim == 4 - num_nodes = expert_load.shape[2] - if node_world_size is None: - assert num_ranks % num_nodes == 0 - node_world_size = num_ranks // num_nodes - assert 0 < node_world_size <= num_ranks and num_ranks % node_world_size == 0 - assert num_nodes == num_ranks // node_world_size - return node_world_size - - -def _expert_locations(redundant_expert_ids: torch.Tensor, num_logical_experts: int) -> torch.Tensor: - """Return ``[layer, logical expert, rank]`` physical-copy occupancy.""" - num_layers, num_ranks, num_redundant_experts_per_rank = redundant_expert_ids.shape - assert num_logical_experts % num_ranks == 0 - num_experts_per_rank = num_logical_experts // num_ranks - locations = torch.zeros( - (num_layers, num_logical_experts, num_ranks), - dtype=torch.bool, - device=redundant_expert_ids.device, - ) - expert_ids = torch.arange(num_logical_experts, device=locations.device) - owners = expert_ids // num_experts_per_rank - locations[:, expert_ids, owners] = True - layers = torch.arange(num_layers, device=locations.device)[:, None] - ranks = torch.arange(num_ranks, device=locations.device).repeat_interleave(num_redundant_experts_per_rank)[None, :] - redundant_ids = redundant_expert_ids.reshape(num_layers, -1) - valid = redundant_ids >= 0 - if torch.any(valid): - expanded_layers = layers.expand_as(redundant_ids) - expanded_ranks = ranks.expand_as(redundant_ids) - locations[ - expanded_layers[valid], - redundant_ids[valid], - expanded_ranks[valid], - ] = True - return locations - - -def _source_route(slots: torch.Tensor, num_nodes: int, node_world_size: int) -> torch.Tensor: - """Route each source node to its local copies, or all copies as fallback.""" - num_ranks = slots.shape[-1] - assert num_ranks % node_world_size == 0 and num_nodes == num_ranks // node_world_size - rank_nodes = torch.arange(num_ranks, device=slots.device) // node_world_size - source_nodes = torch.arange(num_nodes, device=slots.device) - copies = slots.unsqueeze(-3).expand(*slots.shape[:-2], num_nodes, *slots.shape[-2:]) - rank_node_shape = (1,) * slots.ndim + (num_ranks,) - source_node_shape = (1,) * (slots.ndim - 2) + (num_nodes, 1, 1) - local = copies & (rank_nodes.reshape(rank_node_shape) == source_nodes.reshape(source_node_shape)) - selected = torch.where(local.any(dim=-1, keepdim=True), local, copies) - return selected.to(torch.float64) / selected.sum(dim=-1, keepdim=True) - - -def _expert_rank_load_all( - source_load: torch.Tensor, - locations: torch.Tensor, - num_nodes: int, - node_world_size: int, - expert_alignment: int | None, -) -> torch.Tensor: - """Return aligned ``[samples, layers, expert, rank]`` contributions.""" - route = _source_route(locations, num_nodes, node_world_size) - physical_load = torch.einsum("slne,lner->sler", source_load.to(torch.float64), route) - if expert_alignment is not None: - physical_load = torch.ceil(physical_load / expert_alignment) * expert_alignment - return physical_load diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py deleted file mode 100644 index f66d31a62b..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/expert_parallel_state.py +++ /dev/null @@ -1,56 +0,0 @@ -from contextlib import contextmanager -from contextvars import ContextVar -from dataclasses import dataclass -from typing import Iterator, Optional - -import torch - - -_eplb_model_init_disabled: ContextVar[bool] = ContextVar("eplb_model_init_disabled", default=False) - - -def is_eplb_model_init_disabled() -> bool: - return _eplb_model_init_disabled.get() - - -@contextmanager -def disable_eplb_model_init() -> Iterator[None]: - token = _eplb_model_init_disabled.set(True) - try: - yield - finally: - _eplb_model_init_disabled.reset(token) - - -@dataclass -class EPLBState: - num_redundant_experts_per_rank: int - initial_redundant_expert_ids_by_rank: torch.Tensor - logical_to_physical_map: torch.Tensor - logical_replica_count: torch.Tensor - route_counter: torch.Tensor - recording: bool = False - recorded_sample_count: int = 0 - - def next_sample_index(self) -> int: - if not self.recording: - return 0 - sample_index = self.recorded_sample_count % self.route_counter.shape[0] - self.recorded_sample_count += 1 - return sample_index - - -@dataclass(frozen=True) -class ExpertParallelState: - num_logical_experts: int - world_size: int - eplb: Optional[EPLBState] = None - - @property - def num_primary_experts_per_rank(self) -> int: - return self.num_logical_experts // self.world_size - - @property - def num_total_physical_experts(self) -> int: - num_redundant_experts_per_rank = 0 if self.eplb is None else self.eplb.num_redundant_experts_per_rank - return self.num_logical_experts + self.world_size * num_redundant_experts_per_rank diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index cbfd109cca..758e09cdfb 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -8,20 +8,11 @@ get_col_slice_mixin, SliceMixinTpl, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import create_fuse_moe_impl -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, - is_eplb_model_init_disabled, -) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import select_fuse_moe_impl from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, - build_logical_to_physical_map, -) from lightllm.common.quantization.quantize_method import QuantizationMethod -from lightllm.utils.envs_utils import get_env_start_args, get_prefill_eplb_step_interval -from lightllm.utils.dist_utils import get_global_world_size, get_global_rank, get_node_world_size +from lightllm.utils.envs_utils import get_redundancy_expert_ids, get_redundancy_expert_num, get_env_start_args +from lightllm.utils.dist_utils import get_global_world_size, get_global_rank from lightllm.utils.log_utils import init_logger logger = init_logger(__name__) @@ -61,20 +52,21 @@ def __init__( self.moe_intermediate_size = moe_intermediate_size self.quant_method = quant_method assert num_fused_shared_experts in [0, 1], "num_fused_shared_experts can only support 0 or 1 now." - args = get_env_start_args() - self.enable_ep_moe = args.enable_ep_moe + self.enable_ep_moe = get_env_start_args().enable_ep_moe self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self._init_config(network_config) - self._init_expert_parallel_state() - self._init_weight_partition() - self.fuse_moe_impl = create_fuse_moe_impl( + self._init_redundancy_expert_params() + self._init_parallel_params() + self.fuse_moe_impl = select_fuse_moe_impl(self.quant_method, self.enable_ep_moe)( n_routed_experts=self.n_routed_experts, num_fused_shared_experts=self.num_fused_shared_experts, routed_scaling_factor=self.routed_scaling_factor, quant_method=self.quant_method, - expert_parallel_state=self.expert_parallel_state, - ep_moe_backend=args.ep_moe_backend, + redundancy_expert_num=self.redundancy_expert_num, + redundancy_expert_ids_tensor=self.redundancy_expert_ids_tensor, + routed_expert_counter_tensor=self.routed_expert_counter_tensor, + auto_update_redundancy_expert=self.auto_update_redundancy_expert, ) self.lock = threading.Lock() self._moe_weight_finalized = False @@ -89,50 +81,16 @@ def _init_config(self, network_config: Dict[str, Any]): self.routed_scaling_factor = network_config.get("routed_scaling_factor", 1.0) self.scoring_func = network_config.get("scoring_func", "softmax") - def _init_expert_parallel_state(self): - args = get_env_start_args() - self.expert_parallel_state: Optional[ExpertParallelState] = None - # Initial placement metadata is used only while loading checkpoint rows. - self._initial_redundant_expert_ids = [] - self._initial_redundant_expert_idx_to_local_idx = {} - eplb = None - if args.enable_prefill_eplb and not is_eplb_model_init_disabled(): - num_redundant_experts_per_rank = args.eplb_num_redundant_experts_per_rank - all_initial_ids = build_initial_redundant_expert_ids( - self.n_routed_experts, - self.global_world_size, - num_redundant_experts_per_rank, - ) - self._initial_redundant_expert_ids = all_initial_ids[self.global_rank_].tolist() - logical_to_physical, logical_replica_count = build_logical_to_physical_map( - all_initial_ids, - self.n_routed_experts, - source_rank=self.global_rank_, - node_world_size=get_node_world_size(), - ) - # route_counter 每次 prefill dispatch 记录一行。初始阶段连续采样 - # step_interval 个 manager step,兼顾micro batch overlap的两次 dispatch,因此容量设为 - # 2 * step_interval。稳定阶段复用该环形缓冲区,但只把当前短采样窗口内 - # 实际记录的最近行传给 planner,不复制整个缓冲区。 - eplb = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=all_initial_ids, - logical_to_physical_map=logical_to_physical.cuda(), - logical_replica_count=logical_replica_count.cuda(), - route_counter=torch.zeros( - (2 * get_prefill_eplb_step_interval(), self.n_routed_experts), - dtype=torch.int64, - device="cuda", - ), - ) - if self.enable_ep_moe: - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=self.n_routed_experts, - world_size=self.global_world_size, - eplb=eplb, - ) + def _init_redundancy_expert_params(self): + self.redundancy_expert_num = get_redundancy_expert_num() + self.redundancy_expert_ids = get_redundancy_expert_ids(self.layer_num_) + self.auto_update_redundancy_expert: bool = get_env_start_args().auto_update_redundancy_expert + self.redundancy_expert_ids_tensor = torch.tensor(self.redundancy_expert_ids, dtype=torch.int64, device="cuda") + self.routed_expert_counter_tensor = torch.zeros((self.n_routed_experts,), dtype=torch.int64, device="cuda") + # TODO: find out the reason of failure of deepep when redundancy_expert_num is 1. + assert self.redundancy_expert_num != 1, "redundancy_expert_num can not be 1 for some unknown hang of deepep." - def _init_weight_partition(self): + def _init_parallel_params(self): if self.enable_ep_moe: self.tp_rank_ = 0 self.tp_world_size_ = 1 @@ -146,26 +104,27 @@ def _init_weight_partition(self): self.split_inter_size = self.moe_intermediate_size // self.tp_world_size_ if self.enable_ep_moe: assert self.num_fused_shared_experts == 0, "num_fused_shared_experts must be 0 when enable_ep_moe" - eplb = self.expert_parallel_state.eplb - num_redundant_experts_per_rank = 0 if eplb is None else eplb.num_redundant_experts_per_rank logger.debug( f"global_rank {self.global_rank_} layerindex {self.layer_num_} " - f"initial_redundant_expert_ids: {self._initial_redundant_expert_ids}" + f"redundancy_expertids: {self.redundancy_expert_ids}" + ) + self.local_n_routed_experts = self.n_routed_experts // self.global_world_size + self.redundancy_expert_num + n_experts_per_rank = self.n_routed_experts // self.global_world_size + start_expert_id = self.global_rank_ * n_experts_per_rank + self.local_expert_ids = ( + list(range(start_expert_id, start_expert_id + n_experts_per_rank)) + self.redundancy_expert_ids ) - num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank - self.local_n_routed_experts = num_primary_experts_per_rank + num_redundant_experts_per_rank - start_expert_id = self.global_rank_ * num_primary_experts_per_rank self.expert_idx_to_local_idx = { - expert_idx: expert_idx - start_expert_id - for expert_idx in range(start_expert_id, start_expert_id + num_primary_experts_per_rank) + expert_idx: expert_idx - start_expert_id for expert_idx in self.local_expert_ids[:n_experts_per_rank] } - self._initial_redundant_expert_idx_to_local_idx = { - redundant_expert_idx: num_primary_experts_per_rank + i - for (i, redundant_expert_idx) in enumerate(self._initial_redundant_expert_ids) + self.redundancy_expert_idx_to_local_idx = { + redundancy_expert_idx: n_experts_per_rank + i + for (i, redundancy_expert_idx) in enumerate(self.redundancy_expert_ids) } else: self.local_expert_ids = list(range(self.n_routed_experts + self.num_fused_shared_experts)) self.expert_idx_to_local_idx = {expert_idx: i for (i, expert_idx) in enumerate(self.local_expert_ids)} + self.rexpert_idx_to_local_idx = {} def experts( self, @@ -209,12 +168,11 @@ def experts_with_topk( infer_state=None, clamp_limit: Optional[float] = None, alloc_tensor_func=torch.empty, - logical_topk_ids: Optional[torch.Tensor] = None, ) -> torch.Tensor: moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) if moe_capture_callback is not None: - moe_capture_callback(logical_topk_ids if logical_topk_ids is not None else topk_ids) - return self.fuse_moe_impl._fused_experts( + moe_capture_callback(topk_ids) + return self.fuse_moe_impl.fused_experts_with_topk( input_tensor=input_tensor, w13=self.w13, w2=self.w2, @@ -373,8 +331,8 @@ def load_hf_weights(self, weights): self._load_e_score_correction_bias(weights) self._load_per_expert_scale(weights) self._load_weight(self.expert_idx_to_local_idx, weights) - if self._initial_redundant_expert_idx_to_local_idx: - self._load_weight(self._initial_redundant_expert_idx_to_local_idx, weights) + if self.redundancy_expert_num > 0: + self._load_weight(self.redundancy_expert_idx_to_local_idx, weights) def verify_load(self): weight_load_ok = all(all(_weight_pack.load_ok) for _weight_pack in self.w1_list + self.w2_list + self.w3_list) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 96e060faca..9b32284a1a 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -1,38 +1,20 @@ from lightllm.common.quantization.quantize_method import QuantizationMethod from .triton_impl import FuseMoeTriton -from .triton_ep_impl import FuseMoeTritonEP from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM from .mxfp4_impl import FuseMoeMXFP4 -from ..expert_parallel_state import ExpertParallelState -def create_fuse_moe_impl( - *, - n_routed_experts: int, - num_fused_shared_experts: int, - routed_scaling_factor: float, - quant_method: QuantizationMethod, - expert_parallel_state: ExpertParallelState | None = None, - ep_moe_backend: str = "auto", -): +def select_fuse_moe_impl(quant_method: QuantizationMethod, enable_ep_moe: bool): if quant_method.method_name == "mxfp4w4a16-b32-marlin": - if expert_parallel_state is not None: + if enable_ep_moe: raise RuntimeError("mxfp4w4a16-b32-marlin does not support enable_ep_moe yet") - impl_cls = FuseMoeMXFP4 - elif expert_parallel_state is not None: - use_triton_ep = ep_moe_backend == "triton" and quant_method.method_name == "fp8w8a8-b128-deepgemm" - impl_cls = FuseMoeTritonEP if use_triton_ep else FuseMoeDeepGEMM - elif quant_method.method_name == "awq_marlin": - impl_cls = FuseMoeMarlin + return FuseMoeMXFP4 + + if enable_ep_moe: + return FuseMoeDeepGEMM + + if quant_method.method_name == "awq_marlin": + return FuseMoeMarlin else: - impl_cls = FuseMoeTriton - kwargs = dict( - n_routed_experts=n_routed_experts, - num_fused_shared_experts=num_fused_shared_experts, - routed_scaling_factor=routed_scaling_factor, - quant_method=quant_method, - ) - if expert_parallel_state is not None: - kwargs["expert_parallel_state"] = expert_parallel_state - return impl_cls(**kwargs) + return FuseMoeTriton diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py index 35e872df10..1e3ad4b196 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/base_impl.py @@ -1,25 +1,53 @@ import torch -from abc import ABC, abstractmethod -from typing import Callable, Optional, Tuple +from abc import abstractmethod +from typing import Callable, Optional from lightllm.common.quantization.quantize_method import ( WeightPack, QuantizationMethod, ) +from lightllm.utils.dist_utils import ( + get_global_rank, + get_global_world_size, +) -class FuseMoeBaseImpl(ABC): +class FuseMoeBaseImpl: def __init__( self, n_routed_experts: int, num_fused_shared_experts: int, routed_scaling_factor: float, quant_method: QuantizationMethod, + redundancy_expert_num: int, + redundancy_expert_ids_tensor: torch.Tensor, + routed_expert_counter_tensor: torch.Tensor, + auto_update_redundancy_expert: bool, ): self.n_routed_experts = n_routed_experts self.num_fused_shared_experts = num_fused_shared_experts self.routed_scaling_factor = routed_scaling_factor self.quant_method = quant_method + self.global_rank_ = get_global_rank() + self.global_world_size_ = get_global_world_size() + self.ep_n_routed_experts = self.n_routed_experts // self.global_world_size_ + self.total_expert_num_contain_redundancy = ( + self.n_routed_experts + redundancy_expert_num * self.global_world_size_ + ) + + # redundancy expert related + self.redundancy_expert_num = redundancy_expert_num + self.redundancy_expert_ids_tensor = redundancy_expert_ids_tensor + self.routed_expert_counter_tensor = routed_expert_counter_tensor + self.auto_update_redundancy_expert = auto_update_redundancy_expert + # workspace for kernel optimization + self.workspace = self.create_workspace() + + @abstractmethod + def create_workspace(self): + pass + + @abstractmethod def __call__( self, input_tensor: torch.Tensor, @@ -39,62 +67,5 @@ def __call__( per_expert_scale: Optional[torch.Tensor] = None, # Qwen3.5 uses this gate to control fused shared expert aggregation weights. shared_expert_gate: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - topk_weights, topk_ids, origin_topk_ids = self._select_experts( - input_tensor=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - top_k=top_k, - renormalize=renormalize, - use_grouped_topk=use_grouped_topk, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - per_expert_scale=per_expert_scale, - shared_expert_gate=shared_expert_gate, - is_prefill=is_prefill, - preserve_logical_ids=moe_capture_callback is not None, - ) - if moe_capture_callback is not None: - moe_capture_callback(origin_topk_ids) - return self._fused_experts( - input_tensor=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_ids=topk_ids, - router_logits=router_logits, - is_prefill=is_prefill, - ) - - @abstractmethod - def _select_experts( - self, - input_tensor: torch.Tensor, - router_logits: torch.Tensor, - correction_bias: Optional[torch.Tensor], - top_k: int, - renormalize: bool, - use_grouped_topk: bool, - topk_group: int, - num_expert_group: int, - scoring_func: str, - per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - preserve_logical_ids: bool = False, - ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - pass - - @abstractmethod - def _fused_experts( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, ) -> torch.Tensor: pass diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 5b4ca2f515..3f83588c05 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -1,7 +1,6 @@ import torch from typing import Optional, Tuple, Any -from .base_impl import FuseMoeBaseImpl -from ..expert_parallel_state import ExpertParallelState +from .triton_impl import FuseMoeTriton from lightllm.distributed import dist_group_manager from lightllm.common.quantization.quantize_method import WeightPack from lightllm.utils.envs_utils import ( @@ -17,15 +16,83 @@ ) from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType +from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair -class FuseMoeDeepGEMM(FuseMoeBaseImpl): - def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): - super().__init__(*args, **kwargs) - self.expert_parallel_state = expert_parallel_state - self.eplb = expert_parallel_state.eplb - self.ep_balance_counters = None - self._primary_weight_pack_cache = {} +class FuseMoeDeepGEMM(FuseMoeTriton): + def _select_experts( + self, + input_tensor: torch.Tensor, + router_logits: torch.Tensor, + correction_bias: Optional[torch.Tensor], + top_k: int, + renormalize: bool, + use_grouped_topk: bool, + topk_group: int, + num_expert_group: int, + scoring_func: str, + per_expert_scale: Optional[torch.Tensor] = None, + shared_expert_gate: Optional[torch.Tensor] = None, + ): + """Select experts and return topk weights and ids.""" + assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" + from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts + + topk_weights, topk_ids = select_experts( + hidden_states=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + use_grouped_topk=use_grouped_topk, + top_k=top_k, + renormalize=renormalize, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + ) + if self.routed_scaling_factor != 1.0: + topk_weights.mul_(self.routed_scaling_factor) + if per_expert_scale is not None: + topk_weights = topk_weights * per_expert_scale[topk_ids.to(torch.long)].to(topk_weights.dtype) + origin_topk_ids = topk_ids + if self.redundancy_expert_num > 0: + # 因为 redundancy_topk_ids_repair 会修改 topk_ids,所以需要先复制一份 + origin_topk_ids = topk_ids.clone() + redundancy_topk_ids_repair( + topk_ids=topk_ids, + redundancy_expert_ids=self.redundancy_expert_ids_tensor, + ep_expert_num=self.ep_n_routed_experts, + global_rank=self.global_rank_, + expert_counter=self.routed_expert_counter_tensor, + enable_counter=self.auto_update_redundancy_expert, + ) + return topk_weights, topk_ids, origin_topk_ids + + def _fused_experts( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + router_logits: Optional[torch.Tensor] = None, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, + ): + output = fused_experts( + hidden_states=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_idx=topk_ids.to(torch.long), + num_experts=self.total_expert_num_contain_redundancy, # number of all experts contain redundancy + quant_method=self.quant_method, + is_prefill=is_prefill, + previous_event=None, # for overlap + clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, + ) + return output def low_latency_dispatch( self, @@ -49,7 +116,6 @@ def low_latency_dispatch( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, - is_prefill=False, ) return self.low_latency_dispatch_with_topk( @@ -71,8 +137,7 @@ def low_latency_dispatch_with_topk( topk_idx=topk_idx, x=hidden_states, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - # decode 与 EPLB 的物理冗余行刻意隔离:DeepEP 使用原始 logical expert ID。 - num_experts=self.n_routed_experts, + num_experts=self.total_expert_num_contain_redundancy, use_fp8=use_fp8_w8a8, async_finish=False, return_recv_hook=True, @@ -105,7 +170,6 @@ def select_experts_and_quant_input( topk_group=topk_group, num_expert_group=n_group, scoring_func=scoring_func, - is_prefill=True, ) qinput_tensor = quantize_fused_experts_input(hidden_states, w13, self.quant_method) return topk_weights, topk_idx.to(torch.long), qinput_tensor @@ -123,7 +187,7 @@ def dispatch( qinput_tensor, topk_idx=topk_idx, topk_weights=topk_weights, - num_experts=self.expert_parallel_state.num_total_physical_experts, + num_experts=self.total_expert_num_contain_redundancy, num_max_tokens_per_rank=num_max_tokens_per_rank, expert_alignment=128, num_sms=get_ep_num_sms(), @@ -136,23 +200,8 @@ def dispatch( use_tma_aligned_col_major_sf=True, ) - counters = self.ep_balance_counters - if counters is None: - - def hook(): - event.current_stream_wait() - - else: - # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. - route_load = topk_idx.numel() - compute_load = recv_x[0].shape[0] - - def hook(): - event.current_stream_wait() - counters.accumulate( - route_load=route_load, - compute_load=compute_load, - ) + def hook(): + event.current_stream_wait() return recv_x, recv_topk_idx, recv_topk_weights, handle.num_recv_tokens_per_expert_list, handle, hook @@ -166,7 +215,6 @@ def masked_group_gemm( expected_m: int, clamp_limit: Optional[float] = None, ): - w13, w2 = self._primary_weight_pack(w13), self._primary_weight_pack(w2) w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale return masked_group_gemm( @@ -267,118 +315,3 @@ def hook(): event.current_stream_wait() return combined_x, hook - - def _select_experts( - self, - input_tensor: torch.Tensor, - router_logits: torch.Tensor, - correction_bias: Optional[torch.Tensor], - top_k: int, - renormalize: bool, - use_grouped_topk: bool, - topk_group: int, - num_expert_group: int, - scoring_func: str, - per_expert_scale: Optional[torch.Tensor] = None, - shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - preserve_logical_ids: bool = False, - ): - """选择 expert;EPLB prefill 统一由融合路径返回 physical ID。""" - assert shared_expert_gate is None, "fused shared expert as MoE is not supported by DeepGEMM fused MoE" - eplb = self.eplb - if is_prefill is True and eplb is not None: - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - - group_score_topk_num = 2 if topk_group == 4 and num_expert_group == 8 and top_k == 8 else 1 - topk_weights, topk_ids, logical_topk_ids = triton_grouped_topk_eplb( - gating_output=router_logits, - correction_bias=correction_bias, - topk=top_k, - renormalize=renormalize, - num_expert_group=num_expert_group, - topk_group=topk_group, - scoring_func=scoring_func, - logical_to_physical_map=eplb.logical_to_physical_map, - logical_replica_count=eplb.logical_replica_count, - expert_counter=eplb.route_counter, - sample_index=eplb.next_sample_index(), - record_load=eplb.recording, - use_grouped_topk=use_grouped_topk, - return_logical_ids=preserve_logical_ids, - group_score_used_topk_num=group_score_topk_num, - ) - origin_topk_ids = logical_topk_ids if logical_topk_ids is not None else topk_ids - else: - from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts - - topk_weights, topk_ids = select_experts( - hidden_states=input_tensor, - router_logits=router_logits, - correction_bias=correction_bias, - use_grouped_topk=use_grouped_topk, - top_k=top_k, - renormalize=renormalize, - topk_group=topk_group, - num_expert_group=num_expert_group, - scoring_func=scoring_func, - ) - origin_topk_ids = topk_ids - if self.routed_scaling_factor != 1.0: - topk_weights.mul_(self.routed_scaling_factor) - return topk_weights, topk_ids, origin_topk_ids - - def _fused_experts( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, - alloc_tensor_func=torch.empty, - ): - if is_prefill is False: - w13 = self._primary_weight_pack(w13) - w2 = self._primary_weight_pack(w2) - num_experts = self.n_routed_experts - else: - num_experts = self.expert_parallel_state.num_total_physical_experts - - output = fused_experts( - hidden_states=input_tensor, - w13=w13, - w2=w2, - topk_weights=topk_weights, - topk_idx=topk_ids.to(torch.long), - num_experts=num_experts, - quant_method=self.quant_method, - is_prefill=is_prefill, - previous_event=None, # for overlap - clamp_limit=clamp_limit, - alloc_tensor_func=alloc_tensor_func, - ep_balance_counters=self.ep_balance_counters, - ) - return output - - def _primary_weight_pack(self, weight_pack: WeightPack) -> WeightPack: - """返回所有 decode 路径使用的缓存本地主副本视图。""" - if self.eplb is None: - return weight_pack - cache = self._primary_weight_pack_cache - cache_key = id(weight_pack) - primary = cache.get(cache_key) - if primary is None: - num_primary_experts_per_rank = self.expert_parallel_state.num_primary_experts_per_rank - primary = WeightPack( - weight=weight_pack.weight[:num_primary_experts_per_rank], - weight_scale=( - weight_pack.weight_scale[:num_primary_experts_per_rank] - if weight_pack.weight_scale is not None - else None - ), - ) - cache[cache_key] = primary - return primary diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 59ba9fe766..a30a669c18 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -11,10 +11,6 @@ class FuseMoeMarlin(FuseMoeTriton): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.workspace = self.create_workspace() - def create_workspace(self): from lightllm.utils.vllm_utils import HAS_VLLM diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py deleted file mode 100644 index edfc95a8a0..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_ep_impl.py +++ /dev/null @@ -1,42 +0,0 @@ -from typing import Optional - -import torch - -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - ExpertParallelState, -) -from lightllm.common.quantization.quantize_method import WeightPack -from lightllm.distributed import dist_group_manager - -from .triton_impl import FuseMoeTriton - - -class FuseMoeTritonEP(FuseMoeTriton): - """Triton MoE backend for expert-parallel symmetric-memory execution.""" - - def __init__(self, *args, expert_parallel_state: ExpertParallelState, **kwargs): - super().__init__(*args, **kwargs) - self.expert_parallel_state = expert_parallel_state - - def _fused_experts( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, - alloc_tensor_func=torch.empty, - ) -> torch.Tensor: - buffer = dist_group_manager.ep_triton_moe_buffer - return buffer.forward( - input_tensor, - (w13.weight, w13.weight_scale), - (w2.weight, w2.weight_scale), - topk_weights, - topk_ids.to(torch.long), - float("inf") if clamp_limit is None else clamp_limit, - alloc_tensor_func=alloc_tensor_func, - ) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index 43ea96865e..a44bf16c54 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -1,10 +1,36 @@ import torch -from typing import Optional +from typing import Callable, Optional from lightllm.common.quantization.no_quant import WeightPack +from lightllm.common.quantization.quantize_method import QuantizationMethod from .base_impl import FuseMoeBaseImpl class FuseMoeTriton(FuseMoeBaseImpl): + def __init__( + self, + n_routed_experts: int, + num_fused_shared_experts: int, + routed_scaling_factor: float, + quant_method: QuantizationMethod, + redundancy_expert_num: int, + redundancy_expert_ids_tensor: torch.Tensor, + routed_expert_counter_tensor: torch.Tensor, + auto_update_redundancy_expert: bool, + ): + super().__init__( + n_routed_experts=n_routed_experts, + num_fused_shared_experts=num_fused_shared_experts, + routed_scaling_factor=routed_scaling_factor, + quant_method=quant_method, + redundancy_expert_num=redundancy_expert_num, + redundancy_expert_ids_tensor=redundancy_expert_ids_tensor, + routed_expert_counter_tensor=routed_expert_counter_tensor, + auto_update_redundancy_expert=auto_update_redundancy_expert, + ) + + def create_workspace(self): + return None + def _select_experts( self, input_tensor: torch.Tensor, @@ -18,8 +44,6 @@ def _select_experts( scoring_func: str, per_expert_scale: Optional[torch.Tensor] = None, shared_expert_gate: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - preserve_logical_ids: bool = False, ): """Select experts and return topk weights and ids.""" from lightllm.common.basemodel.triton_kernel.fused_moe.topk_select import select_experts @@ -85,3 +109,72 @@ def _fused_experts( limit=clamp_limit, ) return input_tensor + + def fused_experts_with_topk( + self, + input_tensor: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + is_prefill: Optional[bool] = None, + clamp_limit: Optional[float] = None, + alloc_tensor_func=torch.empty, + ): + return self._fused_experts( + input_tensor=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + is_prefill=is_prefill, + clamp_limit=clamp_limit, + alloc_tensor_func=alloc_tensor_func, + ) + + def __call__( + self, + input_tensor: torch.Tensor, + router_logits: torch.Tensor, + w13: WeightPack, + w2: WeightPack, + correction_bias: Optional[torch.Tensor], + scoring_func: str, + top_k: int, + renormalize: bool, + use_grouped_topk: bool, + topk_group: int, + num_expert_group: int, + is_prefill: Optional[bool] = None, + # Callback to capture MoE topk expert ids (routed experts metadata). + moe_capture_callback: Optional[Callable[[torch.Tensor], None]] = None, + per_expert_scale: Optional[torch.Tensor] = None, + shared_expert_gate: Optional[torch.Tensor] = None, + ): + topk_weights, topk_ids, origin_topk_ids = self._select_experts( + input_tensor=input_tensor, + router_logits=router_logits, + correction_bias=correction_bias, + top_k=top_k, + renormalize=renormalize, + use_grouped_topk=use_grouped_topk, + topk_group=topk_group, + num_expert_group=num_expert_group, + scoring_func=scoring_func, + per_expert_scale=per_expert_scale, + shared_expert_gate=shared_expert_gate, + ) + + if moe_capture_callback is not None: + moe_capture_callback(origin_topk_ids) + + output = self._fused_experts( + input_tensor=input_tensor, + w13=w13, + w2=w2, + topk_weights=topk_weights, + topk_ids=topk_ids, + router_logits=router_logits, + is_prefill=is_prefill, + ) + return output diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 7c0e393633..155797b1c5 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -11,7 +11,6 @@ silu_and_mul_masked_post_quant_fwd, ) from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import ( - lightllm_per_token_group_quant_fp8, per_token_group_quant_fp8, ) from lightllm.common.basemodel.triton_kernel.fused_moe.deepep_expanded_layout_kernels import ( @@ -20,7 +19,6 @@ ep_gather_chunk, ep_zero_padding, ) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters from lightllm.utils.envs_utils import ( get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, @@ -31,6 +29,7 @@ from lightllm.utils.tensor_buffer_manager import TensorBufferManager logger = init_logger(__name__) +_MEGA_MOE_STATES: Dict[Tuple[int, int, int, int], Dict[str, Any]] = {} SUPPORTED_EP_EXPERT_DTYPES = ("fp8w8a8-b128-deepgemm", "fp4fp8-b32-deepgemm") @@ -48,8 +47,8 @@ def get_ep_num_sms() -> int: return getattr(dist_group_manager, "ep_num_sms", None) or 0 -def use_mega_moe(quant_method: Any) -> bool: - return getattr(quant_method, "mega_moe_mma_type", None) is not None +def use_sm100_mega_moe(quant_method: Any) -> bool: + return is_sm100_gpu() and quant_method.method_name == "fp4fp8-b32-deepgemm" def check_ep_expert_dtype(quant_method: Any): @@ -98,28 +97,23 @@ def masked_group_gemm( return gemm_out_b -def _get_mega_moe_weights(w13: Any, w2: Any, mma_type: str): - weights = getattr(w13, "_mega_moe_weights", None) - if weights is not None: - return weights - transform_kwargs = {"mma_type": mma_type} if mma_type == "fp8xfp8" else {} - weights = deep_gemm.transform_weights_for_mega_moe( - (w13.weight, w13.weight_scale), - (w2.weight, w2.weight_scale), - **transform_kwargs, +def _get_mega_moe_cache_state(w13: Any, w2: Any): + state_key = ( + w13.weight.data_ptr(), + w13.weight_scale.data_ptr(), + w2.weight.data_ptr(), + w2.weight_scale.data_ptr(), ) - if mma_type == "fp8xfp8": - # Keep the transformed layout in the preallocated weight storage so we do not retain a second - # full copy of the expert weights. Skip copy_ when DeepGEMM already returned an alias. - for target, transformed in zip( - (w13.weight, w13.weight_scale, w2.weight, w2.weight_scale), - (*weights[0], *weights[1]), - ): - if target.data_ptr() != transformed.data_ptr(): - target.copy_(transformed) - weights = ((w13.weight, w13.weight_scale), (w2.weight, w2.weight_scale)) - w13._mega_moe_weights = weights - return weights + return _MEGA_MOE_STATES.setdefault(state_key, {}) + + +def _get_mega_moe_weights(w13: Any, w2: Any, state: Dict[str, Any]): + if "weight_cache" not in state: + state["weight_cache"] = deep_gemm.transform_weights_for_mega_moe( + (w13.weight, w13.weight_scale), + (w2.weight, w2.weight_scale), + ) + return state["weight_cache"] def _get_mega_moe_cumulative_stats(num_local_experts: int, device: torch.device, state: Dict[str, Any]): @@ -130,19 +124,6 @@ def _get_mega_moe_cumulative_stats(num_local_experts: int, device: torch.device, return stats -def prepare_ep_moe_weights(w13: Any, w2: Any, quant_method: Any): - if dist_group_manager.ep_triton_moe_quant_method == quant_method.method_name: - quant_method.mega_moe_mma_type = None - return - mma_type = dist_group_manager.ep_mega_moe_mma_type - if dist_group_manager.ep_mega_moe_quant_method != quant_method.method_name: - quant_method.mega_moe_mma_type = None - return - quant_method.mega_moe_mma_type = mma_type - if mma_type == "fp8xfp8": - _get_mega_moe_weights(w13, w2, mma_type) - - def mega_moe_impl( hidden_states: torch.Tensor, w13: Any, @@ -150,17 +131,15 @@ def mega_moe_impl( topk_weights: torch.Tensor, topk_ids: torch.Tensor, quant_method: Any, - mma_type: str, - clamp_limit: Optional[float] = None, - alloc_tensor_func: Callable = torch.empty, ): - kernel_name = "fp8_fp8_mega_moe" if mma_type == "fp8xfp8" else "fp8_fp4_mega_moe" - if not (HAS_DEEPGEMM and hasattr(deep_gemm, kernel_name)): - raise RuntimeError(f"deep_gemm does not provide {kernel_name} Mega MoE kernel") + if not (HAS_DEEPGEMM and hasattr(deep_gemm, "fp8_fp4_mega_moe")): + raise RuntimeError("deep_gemm does not provide fp8-fp4 Mega MoE kernel") + + from deep_gemm.utils import per_token_cast_to_fp8 buffer = getattr(dist_group_manager, "ep_mega_moe_buffer", None) if buffer is None: - raise RuntimeError("Mega MoE requires dist_group_manager.ep_mega_moe_buffer to be initialized") + raise RuntimeError("SM100 Mega MoE requires dist_group_manager.ep_mega_moe_buffer to be initialized") num_tokens = hidden_states.shape[0] if num_tokens > buffer.num_max_tokens_per_rank: @@ -168,48 +147,27 @@ def mega_moe_impl( f"Mega MoE got {num_tokens} tokens, exceeding num_max_tokens_per_rank={buffer.num_max_tokens_per_rank}" ) - if mma_type == "fp8xfp8": - lightllm_per_token_group_quant_fp8( - x=hidden_states, - group_size=quant_method.block_size, - x_q=buffer.x[:num_tokens], - x_s=buffer.x_sf[:num_tokens], - eps=1e-4, - dtype=buffer.x.dtype, - topk_ids=topk_ids, - topk_weights=topk_weights, - topk_ids_out=buffer.topk_idx[:num_tokens], - topk_weights_out=buffer.topk_weights[:num_tokens], - ) - else: - from deep_gemm.utils import per_token_cast_to_fp8 - - qinput_tensor = per_token_cast_to_fp8( - hidden_states, - use_ue8m0=True, - gran_k=quant_method.block_size, - use_packed_ue8m0=True, - ) - buffer.x[:num_tokens].copy_(qinput_tensor[0]) - buffer.x_sf[:num_tokens].copy_(qinput_tensor[1]) - buffer.topk_idx[:num_tokens].copy_(topk_ids) - buffer.topk_weights[:num_tokens].copy_(topk_weights) - - l1_weights, l2_weights = _get_mega_moe_weights(w13, w2, mma_type) - state = getattr(w13, "_mega_moe_state", None) - if state is None: - state = {} - w13._mega_moe_state = state + qinput_tensor = per_token_cast_to_fp8( + hidden_states, + use_ue8m0=True, + gran_k=quant_method.block_size, + use_packed_ue8m0=True, + ) + state = _get_mega_moe_cache_state(w13, w2) + l1_weights, l2_weights = _get_mega_moe_weights(w13, w2, state) stats = _get_mega_moe_cumulative_stats(w13.weight.shape[0], hidden_states.device, state) - output = alloc_tensor_func(hidden_states.shape, device=hidden_states.device, dtype=hidden_states.dtype) - kernel = getattr(deep_gemm, kernel_name) - kernel( + buffer.x[:num_tokens].copy_(qinput_tensor[0]) + buffer.x_sf[:num_tokens].copy_(qinput_tensor[1]) + buffer.topk_idx[:num_tokens].copy_(topk_ids) + buffer.topk_weights[:num_tokens].copy_(topk_weights) + + output = torch.empty_like(hidden_states) + deep_gemm.fp8_fp4_mega_moe( output, l1_weights, l2_weights, buffer, cumulative_local_expert_recv_stats=stats, - activation_clamp=clamp_limit, ) return output @@ -220,6 +178,16 @@ def quantize_fused_experts_input( quant_method: Any, ): check_ep_expert_dtype(quant_method) + if use_sm100_mega_moe(quant_method): + from deep_gemm.utils import per_token_cast_to_fp8 + + return per_token_cast_to_fp8( + hidden_states, + use_ue8m0=True, + gran_k=quant_method.block_size, + use_packed_ue8m0=True, + ) + block_size_k = 0 if w13.weight.ndim == 3: block_size_k = w13.weight.shape[2] // w13.weight_scale.shape[2] @@ -239,22 +207,12 @@ def fused_experts( previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, alloc_tensor_func: Callable = torch.empty, - ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): check_ep_expert_dtype(quant_method) - mma_type = getattr(quant_method, "mega_moe_mma_type", None) - if mma_type is not None: - return mega_moe_impl( - hidden_states, - w13, - w2, - topk_weights, - topk_idx, - quant_method, - mma_type, - clamp_limit=clamp_limit, - alloc_tensor_func=alloc_tensor_func, - ) + if use_sm100_mega_moe(quant_method): + if clamp_limit is not None: + raise RuntimeError("SM100 Mega MoE does not support clamped SwiGLU yet.") + return mega_moe_impl(hidden_states, w13, w2, topk_weights, topk_idx, quant_method) buffer = dist_group_manager.ep_buffer if is_prefill else dist_group_manager.ep_low_latency_buffer return fused_experts_impl( @@ -274,7 +232,6 @@ def fused_experts( previous_event=previous_event, clamp_limit=clamp_limit, alloc_tensor_func=alloc_tensor_func, - ep_balance_counters=ep_balance_counters, ) @@ -295,7 +252,6 @@ def fused_experts_impl( previous_event: Optional[Any] = None, clamp_limit: Optional[float] = None, alloc_tensor_func: Callable = torch.empty, - ep_balance_counters: Optional[PrefillEPBalanceCounters] = None, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -348,12 +304,6 @@ def fused_experts_impl( do_expand=True, use_tma_aligned_col_major_sf=True, ) - if ep_balance_counters is not None: - # Sent routes are globally conserved by all-to-all; recv_x[0] is the 128-aligned expanded compute load. - ep_balance_counters.accumulate( - route_load=topk_idx.numel(), - compute_load=recv_x[0].shape[0], - ) # Dispatch is synchronous in this path. Its FP8 source is no longer # needed once the received tensors have been produced. del qinput_tensor, input_scale diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py index 7651282b2f..fb0323cd4b 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_topk.py @@ -5,14 +5,6 @@ from triton.language.standard import _log2, sum, zeros_like -@triton.jit -def _eplb_replica_index(token_index, logical_id, replica_count): - """Choose a replica with independent phases for a token's top-k experts.""" - token_hash = token_index.to(tl.uint32) * 2654435769 - expert_hash = logical_id.to(tl.uint32) * 2246822519 - return (token_hash + expert_hash) % replica_count.to(tl.uint32) - - @triton.jit def _compare_and_swap(x, x_1, ids, flip, i: tl.core.constexpr, n_dims: tl.core.constexpr): n_outer: tl.core.constexpr = x.numel >> n_dims @@ -210,172 +202,6 @@ def grouped_topk_kernel( return -@triton.jit -def grouped_topk_eplb_kernel( - gating_output_ptr, - gating_output_stride_m, - gating_output_stride_n, - correction_bias_ptr, - out_topk_weights, - out_topk_weights_stride_m, - out_topk_weights_stride_n, - out_topk_ids, - out_topk_ids_stride_m, - out_topk_ids_stride_n, - out_logical_ids, - out_logical_ids_stride_m, - out_logical_ids_stride_n, - logical_to_physical_ptr, - logical_replica_count_ptr, - expert_counter_ptr, - sample_index, - group_num, - group_expert_num, - total_expert_num, - group_topk_num, - IS_SIGMOID: tl.constexpr, - USE_GROUPED_TOPK: tl.constexpr, - HAS_CORRECTION_BIAS: tl.constexpr, - RETURN_LOGICAL_IDS: tl.constexpr, - EXPERT_GROUP_NUM: tl.constexpr, - EXPERT_GROUP_SIZE: tl.constexpr, - TOPK_NUM: tl.constexpr, - TOPK_BLOCK_SIZE: tl.constexpr, - RENORMALIZE: tl.constexpr, - GROUP_SCORE_USED_TOPK_NUM: tl.constexpr, - COUNTER_NUM_EXPERTS: tl.constexpr, - MAP_SLOTS: tl.constexpr, - RECORD_LOAD: tl.constexpr, - SINGLE_TOKEN: tl.constexpr, -): - """Grouped top-k, EPLB accounting, and replica mapping without a global score workspace.""" - token_index = tl.program_id(axis=0) - offs_group = tl.arange(0, EXPERT_GROUP_NUM) - offs_group_v = tl.arange(0, EXPERT_GROUP_SIZE) - logical_ids = offs_group[:, None] * group_expert_num + offs_group_v[None, :] - valid_expert = ( - (offs_group < group_num)[:, None] - & (offs_group_v < group_expert_num)[None, :] - & (logical_ids < total_expert_num) - ) - hidden_states = tl.load( - gating_output_ptr + token_index * gating_output_stride_m + logical_ids * gating_output_stride_n, - mask=valid_expert, - other=-float("inf"), - ).to(tl.float32) - - if IS_SIGMOID: - old_scores = tl.sigmoid(hidden_states) - else: - group_max = tl.max(hidden_states, axis=1) - global_max = tl.max(group_max, axis=0) - numerators = tl.where(valid_expert, tl.exp(hidden_states - global_max), 0.0) - denominator = tl.sum(tl.sum(numerators, axis=1), axis=0) - old_scores = numerators / denominator - - if HAS_CORRECTION_BIAS: - correction_bias = tl.load(correction_bias_ptr + logical_ids, mask=valid_expert, other=0.0) - scores = tl.where(valid_expert, old_scores + correction_bias, -float("inf")) - else: - scores = tl.where(valid_expert, old_scores, -float("inf")) - - if USE_GROUPED_TOPK: - if GROUP_SCORE_USED_TOPK_NUM == 1: - group_value = tl.max(scores, axis=1) - elif GROUP_SCORE_USED_TOPK_NUM == 2: - first_score, first_index = tl.max(scores, axis=1, return_indices=True) - second_score = tl.max( - tl.where(offs_group_v[None, :] == first_index[:, None], -float("inf"), scores), - axis=1, - ) - group_value = first_score + second_score - else: - sorted_group_scores = tl.sort(scores, dim=1, descending=True) - group_value = tl.sum( - tl.where(offs_group_v[None, :] < GROUP_SCORE_USED_TOPK_NUM, sorted_group_scores, 0.0), - axis=1, - ) - - if EXPERT_GROUP_NUM > 1: - sorted_group_value = tl.sort(group_value, descending=True) - else: - sorted_group_value = group_value - group_topk_value = tl.sum(tl.where(offs_group == group_topk_num - 1, sorted_group_value, 0.0)) - candidate_scores = tl.where( - (group_value >= group_topk_value)[:, None] & valid_expert, - scores, - -float("inf"), - ) - else: - candidate_scores = tl.where(valid_expert, old_scores, -float("inf")) - - sort_block_size: tl.constexpr = EXPERT_GROUP_NUM * EXPERT_GROUP_SIZE - flat_offsets = tl.arange(0, sort_block_size) - candidate_scores = tl.reshape(candidate_scores, (sort_block_size,)) - topk_offsets = tl.arange(0, TOPK_BLOCK_SIZE) - selected_weights = tl.zeros((TOPK_BLOCK_SIZE,), tl.float32) - selected_logical_ids = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) - sum_scores = 0.0 - for topk_index in range(TOPK_NUM): - selected_offset = tl.argmax(candidate_scores, axis=0) - selected_group = selected_offset // EXPERT_GROUP_SIZE - selected_group_offset = selected_offset % EXPERT_GROUP_SIZE - selected_logical_id = selected_group * group_expert_num + selected_group_offset - selected_hidden_state = tl.load( - gating_output_ptr + token_index * gating_output_stride_m + selected_logical_id * gating_output_stride_n - ).to(tl.float32) - if IS_SIGMOID: - selected_weight = tl.sigmoid(selected_hidden_state) - else: - selected_weight = tl.exp(selected_hidden_state - global_max) / denominator - sum_scores += selected_weight - topk_lane = topk_offsets == topk_index - selected_weights = tl.where(topk_lane, selected_weight, selected_weights) - selected_logical_ids = tl.where(topk_lane, selected_logical_id, selected_logical_ids) - candidate_scores = tl.where(flat_offsets == selected_offset, -float("inf"), candidate_scores) - - topk_mask = topk_offsets < TOPK_NUM - if RECORD_LOAD: - tl.atomic_add( - expert_counter_ptr + sample_index * COUNTER_NUM_EXPERTS + selected_logical_ids, - 1, - mask=topk_mask, - sem="relaxed", - ) - if SINGLE_TOKEN: - replica_indices = tl.zeros((TOPK_BLOCK_SIZE,), tl.int32) - else: - replica_counts = tl.load( - logical_replica_count_ptr + selected_logical_ids, - mask=topk_mask, - other=1, - ) - replica_indices = _eplb_replica_index(token_index, selected_logical_ids, replica_counts) - selected_physical_ids = tl.load( - logical_to_physical_ptr + selected_logical_ids * MAP_SLOTS + replica_indices, - mask=topk_mask, - other=-1, - ) - if RENORMALIZE: - selected_weights /= sum_scores - tl.store( - out_topk_weights + token_index * out_topk_weights_stride_m + topk_offsets * out_topk_weights_stride_n, - selected_weights, - mask=topk_mask, - ) - tl.store( - out_topk_ids + token_index * out_topk_ids_stride_m + topk_offsets * out_topk_ids_stride_n, - selected_physical_ids, - mask=topk_mask, - ) - if RETURN_LOGICAL_IDS: - tl.store( - out_logical_ids + token_index * out_logical_ids_stride_m + topk_offsets * out_logical_ids_stride_n, - selected_logical_ids, - mask=topk_mask, - ) - - def triton_grouped_topk( hidden_states: torch.Tensor, gating_output: torch.Tensor, @@ -437,80 +263,3 @@ def triton_grouped_topk( num_stages=1, ) return out_topk_weights, out_topk_ids - - -def triton_grouped_topk_eplb( - gating_output: torch.Tensor, - correction_bias: torch.Tensor, - topk: int, - renormalize: bool, - num_expert_group: int, - topk_group: int, - scoring_func: str, - logical_to_physical_map: torch.Tensor, - logical_replica_count: torch.Tensor, - expert_counter: torch.Tensor, - sample_index: int, - record_load: bool, - use_grouped_topk: bool, - return_logical_ids: bool = False, - group_score_used_topk_num: int = 2, -): - """Fused EPLB prefill top-k returning physical IDs and optional logical IDs.""" - token_num, total_expert_num = gating_output.shape - out_topk_weights = torch.empty((token_num, topk), dtype=torch.float32, device=gating_output.device) - out_topk_ids = torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) - out_logical_ids = ( - torch.empty((token_num, topk), dtype=torch.long, device=gating_output.device) if return_logical_ids else None - ) - if token_num == 0: - return out_topk_weights, out_topk_ids, out_logical_ids - if use_grouped_topk: - assert total_expert_num % num_expert_group == 0 - group_num = num_expert_group - group_expert_num = total_expert_num // num_expert_group - group_topk_num = topk_group - else: - group_num = 1 - group_expert_num = total_expert_num - group_topk_num = 1 - expert_group_num = triton.next_power_of_2(group_num) - expert_group_size = triton.next_power_of_2(group_expert_num) - sort_block_size = expert_group_num * expert_group_size - num_warps = min(max(1, sort_block_size // 256), 8) - grouped_topk_eplb_kernel[(token_num,)]( - gating_output, - *gating_output.stride(), - correction_bias, - out_topk_weights, - *out_topk_weights.stride(), - out_topk_ids, - *out_topk_ids.stride(), - out_logical_ids if out_logical_ids is not None else out_topk_ids, - *(out_logical_ids.stride() if out_logical_ids is not None else out_topk_ids.stride()), - logical_to_physical_map, - logical_replica_count, - expert_counter, - sample_index, - group_num=group_num, - group_expert_num=group_expert_num, - total_expert_num=total_expert_num, - group_topk_num=group_topk_num, - IS_SIGMOID=use_grouped_topk and scoring_func == "sigmoid", - USE_GROUPED_TOPK=use_grouped_topk, - HAS_CORRECTION_BIAS=use_grouped_topk and correction_bias is not None, - RETURN_LOGICAL_IDS=return_logical_ids, - EXPERT_GROUP_NUM=expert_group_num, - EXPERT_GROUP_SIZE=expert_group_size, - TOPK_NUM=topk, - TOPK_BLOCK_SIZE=triton.next_power_of_2(topk), - RENORMALIZE=renormalize, - GROUP_SCORE_USED_TOPK_NUM=group_score_used_topk_num, - COUNTER_NUM_EXPERTS=expert_counter.shape[1], - MAP_SLOTS=logical_to_physical_map.shape[1], - RECORD_LOAD=record_load, - SINGLE_TOKEN=token_num == 1, - num_warps=num_warps, - num_stages=1, - ) - return out_topk_weights, out_topk_ids, out_logical_ids diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py b/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py deleted file mode 100644 index a380f0c02e..0000000000 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/sm90_fp8_triton_ep_moe.py +++ /dev/null @@ -1,575 +0,0 @@ -"""Single-node SM90 FP8 expert parallelism over CUDA symmetric memory.""" - -from typing import Callable, Optional, Tuple - -import deep_gemm -import torch -import torch.distributed as dist -import torch.distributed._symmetric_memory as symm -import triton -import triton.language as tl -from frozendict import frozendict - -from lightllm.common.kernel_config import KernelConfigs - - -DEFAULT_CONFIG = { - "publish_rows_large": 16, - "publish_rows_small": 4, - "publish_num_warps": 4, - "histogram_block": 1024, - "histogram_num_warps": 4, - "pull_num_warps": 4, - "activation_programs": 512, - "activation_rows": 16, - "activation_num_warps": 4, - "return_num_warps": 4, - "combine_rows_large": 4, - "combine_rows_small": 1, - "combine_block": 512, - "combine_num_warps": 4, -} - - -class SM90FP8TritonEPMoEKernelConfig(KernelConfigs): - kernel_name = "sm90_fp8_triton_ep_moe" - - @classmethod - def _params( - cls, - hidden_size: int, - intermediate_size: int, - local_experts: int, - topk: int, - world_size: int, - num_max_tokens_per_rank: int, - alignment: int, - ): - return frozendict( - { - "hidden_size": hidden_size, - "intermediate_size": intermediate_size, - "local_experts": local_experts, - "topk": topk, - "world_size": world_size, - "num_max_tokens_per_rank": num_max_tokens_per_rank, - "alignment": alignment, - } - ) - - @classmethod - def try_to_get_best_config(cls, **kwargs) -> dict: - config = cls.get_the_config(cls._params(**kwargs)) - return dict(DEFAULT_CONFIG if config is None else config) - - @classmethod - def get_config_if_available(cls, **kwargs) -> Optional[dict]: - config = cls.get_the_config(cls._params(**kwargs)) - return None if config is None else dict(config) - - @classmethod - def save_config(cls, config: dict, **kwargs) -> None: - cls.store_config(cls._params(**kwargs), config) - - -@triton.jit -def _publish( - X, - IDS, - WEIGHTS, - Q, - SF, - SID, - SW, - ROWS, - COUNTS, - M, - H: tl.constexpr, - K: tl.constexpr, - E: tl.constexpr, - BE: tl.constexpr, - BK: tl.constexpr, - BM: tl.constexpr, -): - rows = tl.program_id(0) * BM + tl.arange(0, BM) - block = tl.program_id(1) - cols = block * 128 + tl.arange(0, 128) - values = tl.load(X + rows[:, None] * H + cols[None, :], rows[:, None] < M, 0).to(tl.float32) - scale = tl.maximum(tl.max(tl.abs(values), 1), 1e-10) / 448.0 - quant = tl.clamp(tl.div_rn(values, scale[:, None]), -448.0, 448.0).to(Q.dtype.element_ty) - tl.store(Q + rows[:, None] * H + cols[None, :], quant, rows[:, None] < M) - tl.store(SF + rows * (H // 128) + block, scale, rows < M) - if block == 0: - slots = tl.arange(0, BK) - offsets = rows[:, None] * K + slots[None, :] - mask = (rows[:, None] < M) & (slots[None, :] < K) - tl.store(SID + offsets, tl.load(IDS + offsets, mask, -1), mask) - tl.store(SW + offsets, tl.load(WEIGHTS + offsets, mask, 0), mask) - if tl.program_id(0) == 0: - tl.store(ROWS, M) - experts = tl.arange(0, BE) - tl.store(COUNTS + experts, 0, experts < E) - - -@triton.jit -def _histogram( - ID_PTRS, - ROW_PTRS, - COUNTS, - CAP: tl.constexpr, - K: tl.constexpr, - E: tl.constexpr, - RANK: tl.constexpr, - BE: tl.constexpr, - BLOCK: tl.constexpr, -): - peer = tl.program_id(1) - ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) - rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) - offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) - expert = tl.load(ids + offsets, offsets < rows * K, -1).to(tl.int32) - RANK * E - valid = (offsets < rows * K) & (expert >= 0) & (expert < E) - counts = tl.histogram(expert, BE, mask=valid) - experts = tl.arange(0, BE) - tl.atomic_add(COUNTS + experts, counts, experts < E, sem="relaxed") - - -@triton.jit -def _prefix(COUNTS, ENDS, CURSOR, E: tl.constexpr, BE: tl.constexpr, ALIGN: tl.constexpr): - experts = tl.arange(0, BE) - counts = tl.load(COUNTS + experts, experts < E, 0) - padded = tl.cdiv(counts, ALIGN) * ALIGN - ends = tl.cumsum(padded) - tl.store(ENDS + experts, ends, experts < E) - tl.store(CURSOR + experts, ends - padded, experts < E) - - -@triton.jit -def _pull( - X_PTRS, - SF_PTRS, - ID_PTRS, - ROW_PTRS, - CURSOR, - X, - SF, - ROW_MAP, - CAP: tl.constexpr, - H: tl.constexpr, - K: tl.constexpr, - E: tl.constexpr, - RANK: tl.constexpr, - R: tl.constexpr, - WORLD: tl.constexpr, - BK: tl.constexpr, -): - token = tl.program_id(0) // WORLD - peer = (tl.program_id(0) % WORLD + RANK) % WORLD - rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) - if token < rows: - ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) - slots = tl.arange(0, BK) - experts = tl.load(ids + token * K + slots, slots < K, -1).to(tl.int32) - RANK * E - local = (slots < K) & (experts >= 0) & (experts < E) - if tl.sum(local.to(tl.int32), 0) > 0: - source_x = tl.load(X_PTRS + peer).to(tl.pointer_type(tl.float8e4nv)) - source_sf = tl.load(SF_PTRS + peer).to(tl.pointer_type(tl.float32)) - cols = tl.arange(0, H) - scale_cols = tl.arange(0, H // 128) - values = tl.load(source_x + token * H + cols) - scales = tl.load(source_sf + token * (H // 128) + scale_cols) - for slot in range(K): - expert = tl.load(ids + token * K + slot).to(tl.int32) - RANK * E - if (expert >= 0) & (expert < E): - row = tl.atomic_add(CURSOR + expert, 1, sem="relaxed") - tl.store(X + row * H + cols, values) - tl.store(SF + row + scale_cols * R, scales) - tl.store(ROW_MAP + (peer * CAP + token) * K + slot, row) - - -@triton.jit -def _activation( - X, - Q, - SF, - ENDS, - E: tl.constexpr, - R: tl.constexpr, - I: tl.constexpr, - LIMIT: tl.constexpr, - BM: tl.constexpr, -): - total = tl.load(ENDS + E - 1) - blocks_n = I // 128 - for tile in range(tl.program_id(0), tl.cdiv(total, BM) * blocks_n, tl.num_programs(0)): - row = (tile // blocks_n) * BM + tl.arange(0, BM) - block = tile % blocks_n - col = block * 128 + tl.arange(0, 128) - offset = row[:, None] * (2 * I) + col[None, :] - gate = tl.load(X + offset, row[:, None] < total, 0).to(tl.float32) - up = tl.load(X + offset + I, row[:, None] < total, 0) - gate = tl.minimum(gate, LIMIT) - up = tl.clamp(up, -LIMIT, LIMIT) - gate = (gate / (1 + tl.exp(-gate))).to(tl.bfloat16) - act = (up * gate).to(tl.bfloat16).to(tl.float32) - scale = tl.maximum(tl.max(tl.abs(act), 1), 1e-10) / 448.0 - quant = tl.clamp(act / scale[:, None], -448.0, 448.0).to(Q.dtype.element_ty) - tl.store(Q + row[:, None] * I + col[None, :], quant, row[:, None] < total) - tl.store(SF + row + block * R, scale, row < total) - - -@triton.jit -def _return( - X, - ROW_MAP, - ID_PTRS, - WEIGHT_PTRS, - ROW_PTRS, - OUT_PTRS, - CAP: tl.constexpr, - H: tl.constexpr, - K: tl.constexpr, - E: tl.constexpr, - RANK: tl.constexpr, - WORLD: tl.constexpr, -): - token = tl.program_id(0) // WORLD - peer = (tl.program_id(0) % WORLD + RANK) % WORLD - rows = tl.load(tl.load(ROW_PTRS + peer).to(tl.pointer_type(tl.int32))) - if token < rows: - ids = tl.load(ID_PTRS + peer).to(tl.pointer_type(tl.int64)) - weights = tl.load(WEIGHT_PTRS + peer).to(tl.pointer_type(tl.float32)) - cols = tl.arange(0, H) - acc = tl.full((H,), 0.0, tl.float32) - first = -1 - for slot in range(K): - expert = tl.load(ids + token * K + slot).to(tl.int32) - RANK * E - if (expert >= 0) & (expert < E): - row = tl.load(ROW_MAP + (peer * CAP + token) * K + slot) - weight = tl.load(weights + token * K + slot) - acc += tl.load(X + row * H + cols).to(tl.float32) * weight - if first < 0: - first = slot - if first >= 0: - output = tl.multiple_of(tl.load(OUT_PTRS + peer).to(tl.pointer_type(tl.bfloat16)), 16) - tl.store(output + (token * K + first) * H + cols, acc.to(tl.bfloat16)) - - -@triton.jit -def _combine( - RETURNED, - IDS, - Y, - M, - H: tl.constexpr, - K: tl.constexpr, - E: tl.constexpr, - BLOCK: tl.constexpr, - BM: tl.constexpr, -): - row = tl.program_id(0) * BM + tl.arange(0, BM) - col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) - acc = tl.full((BM, BLOCK), 0, tl.float32) - for slot in tl.static_range(K): - expert = tl.load(IDS + row * K + slot, row < M, -1) - valid = (row < M) & (expert >= 0) - for earlier in tl.static_range(slot): - previous = tl.load(IDS + row * K + earlier, row < M, -1) - valid &= (previous < 0) | (previous // E != expert // E) - values = tl.load( - RETURNED + (row[:, None] * K + slot) * H + col[None, :], - valid[:, None], - 0, - ).to(tl.float32) - acc += values - tl.store(Y + row[:, None] * H + col[None, :], acc.to(tl.bfloat16), row[:, None] < M) - - -class SM90FP8TritonEPMoEBuffer: - """SM90 FP8 Triton EP MoE 的共享通信与计算缓冲区。 - - 每个 EP group 只创建一个实例,并由所有 MoE 层复用;各层的专家权重仍由 layer - weight 持有。每个 rank 将本地 token、路由结果发布到 symmetric memory,本 rank - 拉取所有发往本地专家的 token,完成两次 grouped GEMM,再把加权后的局部结果写回 - token 所在 rank。当前生产选择器只在单节点 SM90 Prefill 路径启用该类。 - """ - - def __init__( - self, - group, - num_experts: int, - num_max_tokens_per_rank: int, - topk: int, - hidden_size: int, - intermediate_size: int, - alignment: Optional[int] = None, - ): - self.group = group - self.rank = dist.get_rank(group) - self.world = dist.get_world_size(group) - assert num_experts % self.world == 0 - self.num_experts = num_experts - self.local_experts = num_experts // self.world - self.num_max_tokens_per_rank = num_max_tokens_per_rank - self.topk = topk - self.hidden_size = hidden_size - self.intermediate_size = intermediate_size - - # 普通尾块使用 128 对齐;完整 num_max_tokens_per_rank 块若有单独调优配置, - # 可以切到 256 对齐。显式传入 alignment 时则始终使用同一套配置。 - self.alignment = 128 if alignment is None else alignment - assert hidden_size == 2 * intermediate_size - assert hidden_size % 128 == 0 and intermediate_size % 128 == 0 - assert self.alignment in (128, 256) - self.config = SM90FP8TritonEPMoEKernelConfig.try_to_get_best_config( - hidden_size=hidden_size, - intermediate_size=intermediate_size, - local_experts=self.local_experts, - topk=topk, - world_size=self.world, - num_max_tokens_per_rank=num_max_tokens_per_rank, - alignment=self.alignment, - ) - self.full_alignment = self.alignment - self.full_config = self.config - if alignment is None: - full_config = SM90FP8TritonEPMoEKernelConfig.get_config_if_available( - hidden_size=hidden_size, - intermediate_size=intermediate_size, - local_experts=self.local_experts, - topk=topk, - world_size=self.world, - num_max_tokens_per_rank=num_max_tokens_per_rank, - alignment=256, - ) - if full_config is not None: - self.full_alignment = 256 - self.full_config = full_config - - # 每个本地专家的接收段都需要独立补齐,最坏情况下额外占用 - # local_experts * (alignment - 1) 行。 - max_recv_rows = self.world * num_max_tokens_per_rank * topk - workspace_alignment = max(self.alignment, self.full_alignment) - self.workspace_rows = ( - triton.cdiv(max_recv_rows + self.local_experts * (workspace_alignment - 1), workspace_alignment) - * workspace_alignment - ) - self.handles = [] - self.pointers = [] - self.sources = [] - # sources/pointers 的下标是后续 Triton kernel 的固定协议: - # 0: FP8 hidden states,1: 每 128 列一组的量化 scale - # 2: top-k expert id,3: top-k weight,4: 本 rank 实际 token 数 - # 5: 各专家 owner 写回的 BF16 局部结果 - for shape, dtype in ( - ((num_max_tokens_per_rank, hidden_size), torch.float8_e4m3fn), - ((num_max_tokens_per_rank, hidden_size // 128), torch.float32), - ((num_max_tokens_per_rank, topk), torch.int64), - ((num_max_tokens_per_rank, topk), torch.float32), - ((1,), torch.int32), - ((num_max_tokens_per_rank, topk, hidden_size), torch.bfloat16), - ): - tensor = symm.empty(*shape, dtype=dtype, device="cuda") - handle = symm.rendezvous(tensor, group=group) - assert all(pointer % 16 == 0 for pointer in handle.buffer_ptrs) - self.sources.append(tensor) - self.handles.append(handle) - self.pointers.append(torch.tensor(handle.buffer_ptrs, device="cuda", dtype=torch.int64)) - - # counts 记录每个本地专家收到的 route 数;ends 是对齐后的排他结束位置; - # cursor 从各专家段起点开始原子递增;row_map 保存 (peer, token, top-k slot) - # 到本地 expert-major workspace 行号的映射,供结果回传使用。 - self.counts = torch.empty(self.local_experts, device="cuda", dtype=torch.int32) - self.ends = torch.empty_like(self.counts) - self.cursor = torch.empty_like(self.counts) - self.row_map = torch.empty(self.world * num_max_tokens_per_rank * topk, device="cuda", dtype=torch.int32) - - # x/x_scale 保存按本地专家分段后的 FP8 输入;workspace 依次承载 W1 和 W2 - # 的 BF16 输出。W1 完成后 x 已不再使用,且 H=2I,因此其前半空间可原地 - # 复用为量化后的激活输入,避免再分配一份 FP8 activation buffer。 - self.x = torch.empty(self.workspace_rows, hidden_size, device="cuda", dtype=torch.float8_e4m3fn) - self.x_scale = torch.empty(hidden_size // 128, self.workspace_rows, device="cuda", dtype=torch.float32).T - self.workspace = torch.empty(self.workspace_rows, hidden_size, device="cuda", dtype=torch.bfloat16) - self.activation = self.x.view(self.workspace_rows * 2, intermediate_size)[: self.workspace_rows] - self.activation_scale = self.x_scale[:, : intermediate_size // 128] - - def _get_runtime_alignment_and_config(self, rows: int) -> Tuple[int, dict]: - """完整块使用 full-chunk 调优结果,尾块使用更稳妥的基础配置。""" - if rows == self.num_max_tokens_per_rank: - return self.full_alignment, self.full_config - return self.alignment, self.config - - @property - def allocated_bytes(self) -> int: - """返回本 rank 主要通信与工作区 tensor 的字节数,供基准测试统计。""" - tensors = self.sources + [ - self.counts, - self.ends, - self.cursor, - self.row_map, - self.x, - self.x_scale, - self.workspace, - ] - return sum(tensor.numel() * tensor.element_size() for tensor in tensors) - - def forward( - self, - hidden_states: torch.Tensor, - w1: Tuple[torch.Tensor, torch.Tensor], - w2: Tuple[torch.Tensor, torch.Tensor], - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - clamp_limit: float, - alloc_tensor_func: Callable = torch.empty, - ) -> torch.Tensor: - """执行一次 EP MoE,并返回与 hidden_states 同 shape、同 dtype 的结果。""" - rows = hidden_states.shape[0] - assert rows <= self.num_max_tokens_per_rank - assert hidden_states.shape == (rows, self.hidden_size) - assert hidden_states.dtype == torch.bfloat16 and hidden_states.is_contiguous() - assert topk_ids.shape == topk_weights.shape == (rows, self.topk) - assert topk_ids.dtype == torch.int64 and topk_ids.is_contiguous() - assert topk_weights.dtype == torch.float32 and topk_weights.is_contiguous() - assert w1[0].shape == (self.local_experts, 2 * self.intermediate_size, self.hidden_size) - assert w2[0].shape == (self.local_experts, self.hidden_size, self.intermediate_size) - - alignment, config = self._get_runtime_alignment_and_config(rows) - publish_rows = config["publish_rows_large"] if rows >= 16 else config["publish_rows_small"] - - # 1. 量化本地输入并发布输入、scale 和路由元数据。barrier 之后所有 rank - # 才能安全读取彼此的 symmetric-memory source buffers。 - _publish[(max(1, triton.cdiv(rows, publish_rows)), self.hidden_size // 128)]( - hidden_states, - topk_ids, - topk_weights, - *self.sources[:5], - self.counts, - rows, - self.hidden_size, - self.topk, - self.local_experts, - triton.next_power_of_2(self.local_experts), - triton.next_power_of_2(self.topk), - publish_rows, - num_warps=config["publish_num_warps"], - ) - self.handles[0].barrier(channel=0, timeout_ms=10000) - - # 2. 统计所有 peer 发往本地专家的 route,生成对齐后的 expert-major 分段, - # 再从 peer buffer 拉取输入并记录每条 route 在 workspace 中的行号。 - histogram_block = config["histogram_block"] - _histogram[(triton.cdiv(self.num_max_tokens_per_rank * self.topk, histogram_block), self.world)]( - self.pointers[2], - self.pointers[4], - self.counts, - self.num_max_tokens_per_rank, - self.topk, - self.local_experts, - self.rank, - triton.next_power_of_2(self.local_experts), - histogram_block, - num_warps=config["histogram_num_warps"], - ) - _prefix[(1,)]( - self.counts, - self.ends, - self.cursor, - self.local_experts, - triton.next_power_of_2(self.local_experts), - alignment, - ) - _pull[(self.num_max_tokens_per_rank * self.world,)]( - self.pointers[0], - self.pointers[1], - self.pointers[2], - self.pointers[4], - self.cursor, - self.x, - self.x_scale, - self.row_map, - self.num_max_tokens_per_rank, - self.hidden_size, - self.topk, - self.local_experts, - self.rank, - self.workspace_rows, - self.world, - triton.next_power_of_2(self.topk), - num_warps=config["pull_num_warps"], - ) - - # 3. 在本地 expert-major 布局上执行 W1 -> SwiGLU+FP8 quant -> W2。 - # DeepGEMM 的 contiguous-layout alignment 是进程级状态,调用后必须恢复。 - expected_m = triton.cdiv(self.num_max_tokens_per_rank * self.topk, self.local_experts) - previous_alignment = deep_gemm.get_mk_alignment_for_contiguous_layout() - deep_gemm.set_mk_alignment_for_contiguous_layout(alignment) - try: - deep_gemm.m_grouped_fp8_gemm_nt_contiguous( - (self.x, self.x_scale), - w1, - self.workspace, - self.ends, - use_psum_layout=True, - expected_m_for_psum_layout=expected_m, - ) - _activation[(config["activation_programs"],)]( - self.workspace, - self.activation, - self.activation_scale, - self.ends, - self.local_experts, - self.workspace_rows, - self.intermediate_size, - clamp_limit, - config["activation_rows"], - num_warps=config["activation_num_warps"], - ) - deep_gemm.m_grouped_fp8_gemm_nt_contiguous( - (self.activation, self.activation_scale), - w2, - self.workspace, - self.ends, - use_psum_layout=True, - expected_m_for_psum_layout=expected_m, - ) - finally: - deep_gemm.set_mk_alignment_for_contiguous_layout(previous_alignment) - - # 4. 每个专家 owner 按 row_map 找回源 token,在本地先合并并乘 route weight, - # 然后向源 rank 写回一份局部结果。第二次 barrier 保证写回全部可见。 - _return[(self.num_max_tokens_per_rank * self.world,)]( - self.workspace, - self.row_map, - self.pointers[2], - self.pointers[3], - self.pointers[4], - self.pointers[5], - self.num_max_tokens_per_rank, - self.hidden_size, - self.topk, - self.local_experts, - self.rank, - self.world, - num_warps=config["return_num_warps"], - ) - self.handles[5].barrier(channel=0, timeout_ms=10000) - - # 5. 一个 token 可能命中多个 owner rank;源 rank 对各 owner 的局部结果求和。 - output = alloc_tensor_func(hidden_states.shape, device=hidden_states.device, dtype=hidden_states.dtype) - if rows: - combine_rows = config["combine_rows_large"] if rows >= 16 else config["combine_rows_small"] - combine_block = config["combine_block"] - _combine[(triton.cdiv(rows, combine_rows), self.hidden_size // combine_block)]( - self.sources[5], - self.sources[2], - output, - rows, - self.hidden_size, - self.topk, - self.local_experts, - combine_block, - combine_rows, - num_warps=config["combine_num_warps"], - ) - return output diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py index d2f59de480..1c01cbd638 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py @@ -21,6 +21,7 @@ from lightllm.utils.sgl_utils import sgl_ops from typing import Callable, List, Optional, Tuple from lightllm.common.basemodel.triton_kernel.fused_moe.softmax_topk import softmax_topk +from lightllm.common.triton_utils.autotuner import Autotuner def fused_topk( @@ -167,4 +168,12 @@ def select_experts( hidden_states=hidden_states, gating_output=router_logits, topk=top_k, renormalize=renormalize ) + ######################################## warning ################################################## + # here is used to match autotune feature, make topk_ids more random + if Autotuner.is_autotune_warmup(): + rand_gen = torch.Generator(device="cuda") + rand_gen.manual_seed(router_logits.shape[0]) + router_logits = torch.randn(size=router_logits.shape, generator=rand_gen, dtype=torch.float32, device="cuda") + _, topk_ids = torch.topk(router_logits, k=top_k, dim=1) + return topk_weights, topk_ids diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py index 553d80e586..a439a01f9d 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py @@ -29,7 +29,6 @@ def _per_token_group_quant_fp8( y_q_ptr, y_s_ptr, y_stride, - M, N, eps, fp8_min, @@ -37,31 +36,21 @@ def _per_token_group_quant_fp8( xs_n, xs_stride_m, xs_stride_n, - topk_ids_ptr, - topk_weights_ptr, - topk_ids_out_ptr, - topk_weights_out_ptr, - num_topk, - topk_row_stride, - topk_out_row_stride, BLOCK: tl.constexpr, - TOPK_BLOCK: tl.constexpr, NEED_MASK: tl.constexpr, USE_UE8M0_SCALE: tl.constexpr, - COPY_TOPK: tl.constexpr, - GROUPS_PER_CTA: tl.constexpr, ): - g_id = tl.program_id(0) * GROUPS_PER_CTA + tl.arange(0, GROUPS_PER_CTA) - y_ptr += g_id[:, None] * y_stride - y_q_ptr += g_id[:, None] * y_stride + g_id = tl.program_id(0) + y_ptr += g_id * y_stride + y_q_ptr += g_id * y_stride row_id = g_id // xs_n col_id = g_id % xs_n y_s_ptr += row_id * xs_stride_m + col_id * xs_stride_n - cols = tl.arange(0, BLOCK)[None, :] # N <= BLOCK + cols = tl.arange(0, BLOCK) # N <= BLOCK - if NEED_MASK or GROUPS_PER_CTA > 1: - mask = (g_id[:, None] < M) & (cols < N) + if NEED_MASK: + mask = cols < N other = 0.0 else: mask = None @@ -69,25 +58,15 @@ def _per_token_group_quant_fp8( y = tl.load(y_ptr + cols, mask=mask, other=other).to(tl.float32) # Quant - _absmax = tl.max(tl.abs(y), axis=1) + _absmax = tl.max(tl.abs(y)) if USE_UE8M0_SCALE: y_s = _ceil_to_ue8m0(tl.maximum(_absmax, 1.0e-4) / fp8_max) else: y_s = tl.maximum(_absmax, eps) / fp8_max - y_q = tl.clamp(y / y_s[:, None], fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty) + y_q = tl.clamp(y / y_s, fp8_min, fp8_max).to(y_q_ptr.dtype.element_ty) tl.store(y_q_ptr + cols, y_q, mask=mask) - tl.store(y_s_ptr, y_s, mask=g_id < M) - - if COPY_TOPK: - topk_cols = tl.arange(0, TOPK_BLOCK)[None, :] - topk_mask = (g_id[:, None] < M) & (col_id[:, None] == 0) & (topk_cols < num_topk) - topk_offsets = row_id[:, None] * topk_row_stride + topk_cols - topk_out_offsets = row_id[:, None] * topk_out_row_stride + topk_cols - topk_ids = tl.load(topk_ids_ptr + topk_offsets, mask=topk_mask) - topk_weights = tl.load(topk_weights_ptr + topk_offsets, mask=topk_mask) - tl.store(topk_ids_out_ptr + topk_out_offsets, topk_ids, mask=topk_mask) - tl.store(topk_weights_out_ptr + topk_out_offsets, topk_weights, mask=topk_mask) + tl.store(y_s_ptr, y_s) def lightllm_per_token_group_quant_fp8( @@ -98,10 +77,6 @@ def lightllm_per_token_group_quant_fp8( eps: float = 1e-10, dtype: torch.dtype = torch.float8_e4m3fn, use_ue8m0_scales: bool = False, - topk_ids: Optional[torch.Tensor] = None, - topk_weights: Optional[torch.Tensor] = None, - topk_ids_out: Optional[torch.Tensor] = None, - topk_weights_out: Optional[torch.Tensor] = None, ): """group-wise, per-token quantization on input tensor `x`. Args: @@ -128,26 +103,11 @@ def lightllm_per_token_group_quant_fp8( # heuristics for number of warps num_warps = min(max(BLOCK // 256, 1), 8) num_stages = 1 - copy_topk = topk_ids is not None - if copy_topk: - num_topk = topk_ids.shape[-1] - topk_block = triton.next_power_of_2(num_topk) - topk_row_stride = topk_ids.stride(0) - topk_out_row_stride = topk_ids_out.stride(0) - else: - topk_ids = topk_weights = topk_ids_out = topk_weights_out = x - topk_block = 1 - num_topk = topk_row_stride = topk_out_row_stride = 0 - # Large Mega inputs otherwise launch one CTA per 128-element group. - groups_per_cta = 16 if copy_topk and M >= 4096 else 1 - if groups_per_cta > 1: - num_warps = 4 - _per_token_group_quant_fp8[(triton.cdiv(M, groups_per_cta),)]( + _per_token_group_quant_fp8[(M,)]( x, x_q, x_s, group_size, - M, N, eps, fp8_min=fp8_min, @@ -155,19 +115,9 @@ def lightllm_per_token_group_quant_fp8( xs_n=xs_n, xs_stride_m=xs_stride_m, xs_stride_n=xs_stride_n, - topk_ids_ptr=topk_ids, - topk_weights_ptr=topk_weights, - topk_ids_out_ptr=topk_ids_out, - topk_weights_out_ptr=topk_weights_out, - num_topk=num_topk, - topk_row_stride=topk_row_stride, - topk_out_row_stride=topk_out_row_stride, BLOCK=BLOCK, - TOPK_BLOCK=topk_block, NEED_MASK=BLOCK != group_size, USE_UE8M0_SCALE=use_ue8m0_scales, - COPY_TOPK=copy_topk, - GROUPS_PER_CTA=groups_per_cta, num_warps=num_warps, num_stages=num_stages, ) diff --git a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py new file mode 100644 index 0000000000..692332cc04 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py @@ -0,0 +1,111 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit +def _redundancy_topk_ids_repair_kernel( + topk_ids_ptr, + topk_total_num, + ep_expert_num, + redundancy_expert_num, + global_rank, + redundancy_expert_ids_ptr, + expert_counter_ptr, + BLOCK_SIZE: tl.constexpr, + ENABLE_COUNTER: tl.constexpr, +): + block_index = tl.program_id(0) + offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offs_d < topk_total_num + current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) + + if ENABLE_COUNTER: + tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) + + # Remap original expert IDs to a new space that accounts for redundant expert slots. + new_current_topk_ids = (current_topk_ids // ep_expert_num) * redundancy_expert_num + current_topk_ids + + for i in tl.range(0, redundancy_expert_num, step=1, num_stages=3): + cur_redundancy_expert_id = tl.load(redundancy_expert_ids_ptr + i) + cur_redundancy_expert_id = ( + cur_redundancy_expert_id // ep_expert_num + ) * redundancy_expert_num + cur_redundancy_expert_id + new_current_topk_ids = tl.where( + new_current_topk_ids == cur_redundancy_expert_id, + (ep_expert_num + redundancy_expert_num) * (global_rank) + ep_expert_num + i, + new_current_topk_ids, + ) + + tl.store(topk_ids_ptr + offs_d, new_current_topk_ids, mask=mask) + return + + +@torch.no_grad() +def redundancy_topk_ids_repair( + topk_ids: torch.Tensor, + redundancy_expert_ids: torch.Tensor, + ep_expert_num: int, + global_rank: int, + expert_counter: torch.Tensor = None, + enable_counter: bool = False, +): + assert topk_ids.is_contiguous() + assert len(topk_ids.shape) == 2 + assert redundancy_expert_ids is not None + redundancy_expert_num = redundancy_expert_ids.shape[0] + BLOCK_SIZE = 512 + grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) + num_warps = 4 + + _redundancy_topk_ids_repair_kernel[grid]( + topk_ids_ptr=topk_ids, + topk_total_num=topk_ids.numel(), + ep_expert_num=ep_expert_num, + redundancy_expert_num=redundancy_expert_num, + global_rank=global_rank, + redundancy_expert_ids_ptr=redundancy_expert_ids, + expert_counter_ptr=expert_counter, + BLOCK_SIZE=BLOCK_SIZE, + ENABLE_COUNTER=enable_counter, + num_warps=num_warps, + num_stages=3, + ) + return + + +@triton.jit +def _expert_id_counter_kernel( + topk_ids_ptr, + topk_total_num, + expert_counter_ptr, + BLOCK_SIZE: tl.constexpr, +): + block_index = tl.program_id(0) + offs_d = block_index * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offs_d < topk_total_num + current_topk_ids = tl.load(topk_ids_ptr + offs_d, mask=mask, other=0) + tl.atomic_add(expert_counter_ptr + current_topk_ids, 1, mask=mask) + return + + +@torch.no_grad() +def expert_id_counter( + topk_ids: torch.Tensor, + expert_counter: torch.Tensor, +): + assert topk_ids.is_contiguous() + assert len(topk_ids.shape) == 2 + BLOCK_SIZE = 512 + grid = (triton.cdiv(topk_ids.numel(), BLOCK_SIZE),) + num_warps = 4 + + _expert_id_counter_kernel[grid]( + topk_ids_ptr=topk_ids, + topk_total_num=topk_ids.numel(), + expert_counter_ptr=expert_counter, + BLOCK_SIZE=BLOCK_SIZE, + num_warps=num_warps, + num_stages=1, + ) + return diff --git a/lightllm/common/eplb_utils.py b/lightllm/common/eplb_utils.py deleted file mode 100644 index 63e3dea069..0000000000 --- a/lightllm/common/eplb_utils.py +++ /dev/null @@ -1,16 +0,0 @@ -"""Small, dependency-free EPLB helpers shared by transfer and model profiling.""" - - -EPLB_MAX_STAGING_DEPTH = 8 - - -def extract_eplb_expert_tensors(weight): - result = [] - for pack_name in ("w13", "w2"): - pack = getattr(weight, pack_name) - for value_name in ("weight", "weight_scale"): - tensor = getattr(pack, value_name, None) - if tensor is not None: - assert tensor.ndim >= 1 and tensor.is_contiguous(), f"{pack_name}.{value_name} must be contiguous" - result.append((f"{pack_name}.{value_name}", tensor)) - return result diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index 6b306714a6..dcdbb2ecc3 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -425,7 +425,6 @@ def __init__( cpu_cache_token_page_size: int = DSV4_CPU_CACHE_TOKEN_PAGE_SIZE, always_copy=False, mem_fraction=0.9, - memory_reservations=None, ): assert head_num == 1, "DeepSeek-V4 是 MLA(MQA),dense latent 的 head_num 必须为 1" assert head_dim == self.mla_head_dim, f"DeepSeek-V4 packed KV 期望 head_dim={self.mla_head_dim}" @@ -465,16 +464,7 @@ def __init__( self.layer_to_c128_idx[lid] = c128 c128 += 1 - super().__init__( - size, - dtype, - head_num, - head_dim, - layer_num, - always_copy, - mem_fraction, - memory_reservations=memory_reservations, - ) + super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction) # ------------------------------------------------------------------ sizing @staticmethod diff --git a/lightllm/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py index 83dc2937ea..c2668d1c34 100755 --- a/lightllm/common/kv_cache_mem_manager/mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py @@ -31,30 +31,13 @@ class MemoryManager: operator_class = NormalMemOperator - def __init__( - self, - size, - dtype, - head_num, - head_dim, - layer_num, - always_copy=False, - mem_fraction=0.9, - memory_reservations=None, - ): + def __init__(self, size, dtype, head_num, head_dim, layer_num, always_copy=False, mem_fraction=0.9): self.size = size self.head_num = head_num self.head_dim = head_dim self.layer_num = layer_num self.always_copy = always_copy self.dtype = dtype - # Named reservations are allocations made after KV profiling. They are - # deliberately outside get_fixed_memory_size(): model-specific exact KV - # geometry owns fixed bytes, while these values are deducted once from - # the profile budget only. - self.memory_reservations = dict(memory_reservations or {}) - if any(value < 0 for value in self.memory_reservations.values()): - raise ValueError(f"memory reservations must be non-negative: {self.memory_reservations}") # profile the max total token num if the size is None self.profile_size(mem_fraction) @@ -118,16 +101,11 @@ def profile_size(self, mem_fraction): available_memory = get_available_gpu_memory(world_size) - get_total_gpu_memory() * (1 - mem_fraction) cell_size = self.get_cell_size() fixed_memory_size = self.get_fixed_memory_size() - reservations = getattr(self, "memory_reservations", {}) - reserved_memory_size = sum(reservations.values()) pd_kv_move_buffer_size = self.get_pd_kv_move_buffer_size() - available_memory_bytes = ( - available_memory * 1024 ** 3 - fixed_memory_size - reserved_memory_size - pd_kv_move_buffer_size - ) + available_memory_bytes = available_memory * 1024 ** 3 - fixed_memory_size - pd_kv_move_buffer_size if available_memory_bytes <= 0: raise RuntimeError( f"{type(self).__name__} fixed buffers require {fixed_memory_size / 1024**3:.2f} GB, " - f"plus {reserved_memory_size / 1024**3:.2f} GB reservations, " f"plus {pd_kv_move_buffer_size / 1024**3:.2f} GB for the PD KV transfer buffer, " f"but only {available_memory:.2f} GB is available" ) @@ -139,7 +117,6 @@ def profile_size(self, mem_fraction): logger.info( f"{str(available_memory)} GB space is available after load the model weight\n" f"{str(fixed_memory_size / 1024 ** 2)} MB is reserved for fixed KV cache buffers\n" - f"{reservations} bytes are reserved for post-profile model buffers\n" f"{str(pd_kv_move_buffer_size / 1024 ** 2)} MB is reserved for PD KV transfer buffer\n" f"{str(cell_size / 1024 ** 2)} MB is the size of one token kv cache\n" f"{self.size} is the profiled max_total_token_num with the mem_fraction {mem_fraction}\n" diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 85257288f4..5ce122e82e 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -23,18 +23,10 @@ def __init__(self): self.cache_manager = g_cache_manager assert HAS_DEEPGEMM, "deepgemm is not installed, you can't use quant api of it" - self.mega_moe_mma_type = None def quantize(self, weight: torch.Tensor, output: WeightPack): raise NotImplementedError("Not implemented") - def finalize_moe_weight(self, moe_weight): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( - prepare_ep_moe_weights, - ) - - prepare_ep_moe_weights(moe_weight.w13, moe_weight.w2, self) - def apply( self, input_tensor: torch.Tensor, diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index 1376e4ab6a..c6b14e990b 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -24,11 +24,12 @@ from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union from lightllm.utils.log_utils import init_logger +from lightllm.utils.device_utils import has_nvlink from lightllm.utils.envs_utils import ( - enable_env_vars, get_env_start_args, get_deepep_num_max_dispatch_tokens_per_rank_prefill, get_deepep_num_max_dispatch_tokens_per_rank_decode, + get_redundancy_expert_num, ) from lightllm.utils.dist_utils import ( get_global_world_size, @@ -36,17 +37,10 @@ create_new_group_for_current_dp, create_dp_special_inter_group, ) -from lightllm.utils.device_utils import ( - get_device_sm_count, - has_nvlink, - is_sm90_gpu, - is_sm100_gpu, -) +from lightllm.utils.device_utils import get_device_sm_count, is_sm100_gpu from lightllm.utils.torch_dtype_utils import get_torch_dtype logger = init_logger(__name__) -FP8_MOE_QUANT_METHOD = "fp8w8a8-b128-deepgemm" -FP4_MOE_QUANT_METHOD = "fp4fp8-b32-deepgemm" def get_deep_ep_prefill_moe_workspace_size( @@ -164,15 +158,10 @@ class DistributeGroupManager: def __init__(self): self.groups = [] self.dp_control_group = None - self.ep_balance_monitor_group = None self.ep_buffer = None self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None - self.ep_mega_moe_mma_type = None - self.ep_mega_moe_quant_method = None - self.ep_triton_moe_buffer = None - self.ep_triton_moe_quant_method = None self.ep_num_sms = None def __len__(self): @@ -189,14 +178,6 @@ def create_groups(self, group_size: int): self.groups.append(group) if args.dp > 1: self.dp_control_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") - if ( - getattr(args, "enable_ep_moe", False) - and not getattr(args, "disable_ep_balance_monitor", False) - and getattr(args, "run_mode", "normal") != "decode" - and not getattr(args, "enable_prefill_cudagraph", False) - and not is_sm100_gpu() - ): - self.ep_balance_monitor_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") return def get_default_group(self) -> CustomProcessGroup: @@ -240,9 +221,9 @@ def new_deepep_group( """初始化 DeepEP 通信组以及当前模型实际需要的 MoE buffer。 ``expert_quant_method_names`` 是各 MoE 层最终绑定的 quant method 名称集合。 - 同一个模型可能逐层混用多种 expert quant method:满足约束的 SM100 FP4 和 - SM90 FP8 层走 Mega/Triton EP MoE,其他层走 DeepEP legacy 路径。这里只为实际 - 存在的执行路径分配 buffer,避免为未使用的路径长期占用显存。 + 同一个模型可能逐层混用 FP4 和 FP8:SM100 FP4 层走 Mega MoE,其他层走 + DeepEP legacy low-latency 路径。这里只为实际存在的执行路径分配 buffer, + 避免为未使用的路径长期占用显存。 """ args = get_env_start_args() enable_ep_moe = args.enable_ep_moe @@ -255,10 +236,6 @@ def new_deepep_group( self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None - self.ep_mega_moe_mma_type = None - self.ep_mega_moe_quant_method = None - self.ep_triton_moe_buffer = None - self.ep_triton_moe_quant_method = None self.ep_num_sms = None return assert HAS_DEEPEP, "deep_ep is required for expert parallelism" @@ -275,15 +252,7 @@ def new_deepep_group( self.ll_num_tokens = prefill_num_max_dispatch_tokens_per_rank self.ll_decode_num_tokens = decode_num_max_dispatch_tokens_per_rank self.ll_hidden = hidden_size - total_redundant_experts = ( - get_env_start_args().eplb_num_redundant_experts_per_rank * global_world_size - if get_env_start_args().enable_prefill_eplb - else 0 - ) - self.ll_prefill_num_experts = n_routed_experts + total_redundant_experts - # EPLB's redundant rows are a prefill-only physical layout; decode - # always routes the logical expert space. - self.ll_decode_num_experts = n_routed_experts + self.ll_num_experts = n_routed_experts + get_redundancy_expert_num() * global_world_size self.ep_buffer = deep_ep.ElasticBuffer( deepep_group, num_max_tokens_per_rank=self.ll_num_tokens, @@ -293,43 +262,31 @@ def new_deepep_group( allow_multiple_reduction=True, ) self.ep_mega_moe_buffer = None - self.ep_triton_moe_buffer = None self.ep_low_latency_buffer = None self.ep_prefill_moe_workspace = None if not expert_quant_method_names: raise ValueError("No valid MoE quant method was found while initializing DeepEP buffers") - self.ep_mega_moe_mma_type = None - self.ep_mega_moe_quant_method = None - self.ep_triton_moe_quant_method = None - if args.ep_moe_backend == "triton" and FP8_MOE_QUANT_METHOD in expert_quant_method_names: - self.ep_triton_moe_quant_method = FP8_MOE_QUANT_METHOD - - if ( - self.ep_triton_moe_quant_method is None - and is_sm100_gpu() - and FP4_MOE_QUANT_METHOD in expert_quant_method_names - ): - self.ep_mega_moe_mma_type = "fp8xfp4" - self.ep_mega_moe_quant_method = FP4_MOE_QUANT_METHOD - elif self.ep_triton_moe_quant_method is None and ( - enable_env_vars("LIGHTLLM_ENABLE_SM90_FP8_MEGA_MOE") - and is_sm90_gpu() - and FP8_MOE_QUANT_METHOD in expert_quant_method_names - and total_redundant_experts == 0 - and args.nnodes == 1 - and not args.enable_rl - ): - self.ep_mega_moe_mma_type = "fp8xfp8" - self.ep_mega_moe_quant_method = FP8_MOE_QUANT_METHOD - - enable_mega_moe_buffer = self.ep_mega_moe_mma_type is not None - enable_triton_ep_moe_buffer = self.ep_triton_moe_quant_method is not None - fused_moe_quant_method = self.ep_triton_moe_quant_method or self.ep_mega_moe_quant_method - has_legacy_moe_layer = fused_moe_quant_method is None or any( - method_name != fused_moe_quant_method for method_name in expert_quant_method_names - ) + mega_moe_quant_method = "fp4fp8-b32-deepgemm" + is_sm100 = is_sm100_gpu() + + # Buffer 选择规则: + # 1. 非 SM100 不支持 Mega MoE,只初始化 legacy low-latency buffer; + # 2. SM100 全部 MoE 层为 FP4,只初始化 Mega MoE buffer; + # 3. SM100 全部 MoE 层为 FP8,只初始化 legacy low-latency buffer; + # 4. SM100 逐层混合 FP4/FP8,两套 buffer 都要初始化。 + if is_sm100: + # 只要存在一个 FP4 MoE 层,就需要 Mega MoE buffer;只要存在一个非 FP4 + # MoE 层,就需要 legacy low-latency buffer。FP4/FP8 逐层混用时两者都会初始化。 + has_mega_moe_layer = mega_moe_quant_method in expert_quant_method_names + has_legacy_moe_layer = any( + method_name != mega_moe_quant_method for method_name in expert_quant_method_names + ) + enable_mega_moe_buffer = has_mega_moe_layer + else: + enable_mega_moe_buffer = False + has_legacy_moe_layer = True enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" @@ -338,10 +295,7 @@ def new_deepep_group( # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 # 空闲的本地 RDMA storage 复用为分块 grouped GEMM 的临时 workspace。 decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( - self.ll_decode_num_tokens, - self.ll_hidden, - global_world_size, - self.ll_decode_num_experts, + self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts ) num_rdma_bytes = decode_size_hint # normal 节点同时执行 Prefill 和 Decode,复用的 RDMA buffer 必须覆盖全部 Prefill workspace。 @@ -351,7 +305,7 @@ def new_deepep_group( hidden_size=self.ll_hidden, intermediate_size=moe_intermediate_size, num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_prefill_num_experts, + num_experts=self.ll_num_experts, world_size=global_world_size, hidden_dtype=get_torch_dtype(args.data_type), ) @@ -360,7 +314,7 @@ def new_deepep_group( deepep_group, num_rdma_bytes=num_rdma_bytes, low_latency_mode=True, - num_qps_per_rank=(self.ll_decode_num_experts // global_world_size), + num_qps_per_rank=(self.ll_num_experts // global_world_size), ) self.ep_prefill_moe_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( torch.uint8, use_rdma_buffer=True @@ -372,7 +326,7 @@ def new_deepep_group( hidden_size=self.ll_hidden, intermediate_size=moe_intermediate_size, num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_prefill_num_experts, + num_experts=self.ll_num_experts, world_size=global_world_size, hidden_dtype=get_torch_dtype(args.data_type), ) @@ -382,64 +336,34 @@ def new_deepep_group( device=torch.device("cuda", torch.cuda.current_device()), ) - theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_prefill_num_experts, num_experts_per_tok) - use_all_sms_for_fp8 = enable_triton_ep_moe_buffer or self.ep_mega_moe_mma_type == "fp8xfp8" - deepep_sms = 0 if use_all_sms_for_fp8 and not has_legacy_moe_layer else theoretical_sms - low_latency_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_decode_num_experts, num_experts_per_tok) - self._set_num_sms_for_deep_gemm(deepep_sms, low_latency_sms) - - if enable_triton_ep_moe_buffer: - from lightllm.common.basemodel.triton_kernel.fused_moe.sm90_fp8_triton_ep_moe import ( - SM90FP8TritonEPMoEBuffer, - ) - - self.ep_triton_moe_buffer = SM90FP8TritonEPMoEBuffer( - deepep_group, - num_experts=self.ll_decode_num_experts, - num_max_tokens_per_rank=self.ll_num_tokens, - topk=num_experts_per_tok, - hidden_size=self.ll_hidden, - intermediate_size=moe_intermediate_size, - ) - if enable_mega_moe_buffer: + # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy + # low-latency buffer,因此纯 FP4 模型无需承担后者的大块 RDMA 显存。 if moe_intermediate_size is None: - raise ValueError("Mega MoE requires moe_intermediate_size or intermediate_size in model config") + raise ValueError("SM100 Mega MoE requires moe_intermediate_size or intermediate_size in model config") import deep_gemm - num_max_tokens_per_rank = ( - self.ll_decode_num_tokens - if self.ep_mega_moe_mma_type == "fp8xfp8" and args.run_mode == "decode" - else self.ll_num_tokens - ) - mega_buffer_kwargs = ( - {"mma_type": self.ep_mega_moe_mma_type} if self.ep_mega_moe_mma_type == "fp8xfp8" else {} - ) self.ep_mega_moe_buffer = deep_gemm.get_symm_buffer_for_mega_moe( deepep_group, - self.ll_decode_num_experts, - num_max_tokens_per_rank, + self.ll_num_experts, + self.ll_num_tokens, num_experts_per_tok, self.ll_hidden, moe_intermediate_size, - **mega_buffer_kwargs, ) logger.info( "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " - "mega_moe=%s, mega_moe_mma_type=%s, triton_ep_moe=%s, ll_prefill_num_experts=%s, " - "ll_decode_num_experts=%s, expert_quant_method_names=%s", + "mega_moe=%s, expert_quant_method_names=%s", enable_low_latency_buffer, self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, - self.ep_mega_moe_mma_type, - enable_triton_ep_moe_buffer, - self.ll_prefill_num_experts, - self.ll_decode_num_experts, sorted(expert_quant_method_names), ) + theoretical_sms = self.ep_buffer.get_theoretical_num_sms(self.ll_num_experts, num_experts_per_tok) + self._set_num_sms_for_deep_gemm(theoretical_sms) - def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): + def _set_num_sms_for_deep_gemm(self, deepep_sms: int): try: try: from deep_gemm.jit_kernels.utils import set_num_sms @@ -448,19 +372,11 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int, low_latency_sms: int): device_sms = get_device_sm_count() deepep_sms = max(0, min(deepep_sms, max(device_sms - 2, 0))) - low_latency_sms = max(0, min(low_latency_sms, max(device_sms - 2, 0))) self.ep_num_sms = deepep_sms if self.ep_low_latency_buffer is not None: - # This setting controls the legacy low-latency buffer; keep - # its SM reservation based on decode's logical expert count. - deep_ep.Buffer.set_num_sms(low_latency_sms - low_latency_sms % 2) - deep_gemm_sms = max(device_sms - deepep_sms, 2) - if self.ep_mega_moe_mma_type == "fp8xfp8": - deep_gemm_sms -= deep_gemm_sms % 2 - set_num_sms(deep_gemm_sms) + deep_ep.Buffer.set_num_sms(deepep_sms - deepep_sms % 2) + set_num_sms(max(device_sms - deepep_sms, 2)) except BaseException as e: - if self.ep_mega_moe_mma_type is not None: - raise RuntimeError("Failed to reserve a fixed SM pool before allocating the Mega MoE buffer") from e logger.warning(f"set num sms for deep_gemm failed: {e}") def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: @@ -500,7 +416,7 @@ def clear_deepep_buffer(self): """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( - self.ll_decode_num_tokens, self.ll_hidden, self.ll_decode_num_experts + self.ll_decode_num_tokens, self.ll_hidden, self.ll_num_experts ) diff --git a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py index 5464c3adcb..3254031056 100644 --- a/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek2/layer_infer/transformer_layer_infer.py @@ -7,8 +7,9 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd from lightllm.models.deepseek2.infer_struct import Deepseek2InferStateInfo -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( + use_sm100_mega_moe, +) from functools import partial from lightllm.models.llama.yarn_rotary_utils import get_deepseek_mscale from lightllm.utils.envs_utils import get_env_start_args @@ -299,11 +300,7 @@ def overlap_tpsp_token_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if ( - not self.is_moe - or use_mega_moe(layer_weight.experts.quant_method) - or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -429,11 +426,7 @@ def overlap_tpsp_context_forward( infer_state1: Deepseek2InferStateInfo, layer_weight: Deepseek2TransformerLayerWeight, ): - if ( - not self.is_moe - or use_mega_moe(layer_weight.experts.quant_method) - or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 8c547c35bd..c51ccb0e1c 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -2,10 +2,8 @@ import triton from lightllm.common.basemodel import BaseLayerInfer, TransformerLayerInferTpl from lightllm.common.basemodel.attention.base_att import AttControl -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_sm100_mega_moe from lightllm.common.basemodel.triton_kernel.fused_moe.moe_silu_and_mul import silu_and_mul_fwd -from lightllm.common.basemodel.moe_route_info_manager import get_moe_capture_callback from lightllm.models.deepseek3_2.layer_infer.transformer_layer_infer import Deepseek3_2TransformerLayerInfer from lightllm.models.deepseek_v4.layer_weights.transformer_layer_weight import DeepseekV4TransformerLayerWeight from lightllm.utils.envs_utils import get_env_start_args @@ -151,11 +149,7 @@ def overlap_tpsp_context_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if ( - not self.enable_ep_moe - or use_mega_moe(experts.quant_method) - or isinstance(experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): input_embdings = self.context_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.context_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 @@ -165,7 +159,8 @@ def overlap_tpsp_context_forward( x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) - weights0, indices0, _ = self._select_experts(logits0, infer_state, layer_weight) + weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) + indices0 = indices0.to(torch.long) qinput0 = experts.quantize_dispatch_input(x0) from deep_ep import ElasticBuffer @@ -178,7 +173,7 @@ def overlap_tpsp_context_forward( x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) - weights1, indices1, _ = self._select_experts(logits1, infer_state1, layer_weight) + weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) recv_x0, recv_indices0, recv_weights0, recv_count0, handle0, dispatch_hook0 = experts.dispatch( qinput0, @@ -259,11 +254,7 @@ def overlap_tpsp_token_forward( layer_weight: DeepseekV4TransformerLayerWeight, ): experts = layer_weight.experts_ - if ( - not self.enable_ep_moe - or use_mega_moe(experts.quant_method) - or isinstance(experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.enable_ep_moe or use_sm100_mega_moe(experts.quant_method): input_embdings = self.token_forward(input_embdings, infer_state, layer_weight) input_embdings1 = self.token_forward(input_embdings1, infer_state1, layer_weight) return input_embdings, input_embdings1 @@ -273,7 +264,7 @@ def overlap_tpsp_token_forward( x0, residual0, post_mix0, res_mix0 = self._hc_ffn_in(x0, residual0, post_mix0, res_mix0, layer_weight) x0 = self._tpsp_allgather(x0.view(-1, self.embed_dim_), infer_state) logits0 = layer_weight.gate_weight_.mm(x0, out_dtype=torch.float32) - weights0, indices0, _ = self._select_experts(logits0, infer_state, layer_weight) + weights0, indices0 = self._select_experts(logits0, infer_state, layer_weight) infer_state1.call_overlap_hook() shared0 = self._ffn_tp(x0, infer_state, layer_weight) @@ -286,7 +277,7 @@ def overlap_tpsp_token_forward( x1, residual1, post_mix1, res_mix1 = self._hc_ffn_in(x1, residual1, post_mix1, res_mix1, layer_weight) x1 = self._tpsp_allgather(x1.view(-1, self.embed_dim_), infer_state1) logits1 = layer_weight.gate_weight_.mm(x1, out_dtype=torch.float32) - weights1, indices1, _ = self._select_experts(logits1, infer_state1, layer_weight) + weights1, indices1 = self._select_experts(logits1, infer_state1, layer_weight) dispatch_hook0() shared1 = self._ffn_tp(x1, infer_state1, layer_weight) @@ -540,7 +531,6 @@ def _routed_experts( x, weights, indices, - logical_topk_ids, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight, ): @@ -548,7 +538,6 @@ def _routed_experts( input_tensor=x, topk_weights=weights, topk_ids=indices, - logical_topk_ids=logical_topk_ids, is_prefill=infer_state.is_prefill, infer_state=infer_state, clamp_limit=float(self.swiglu_limit), @@ -571,25 +560,20 @@ def _ffn(self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV x = self._tpsp_allgather(input=x, infer_state=infer_state) logits = layer_weight.gate_weight_.mm(x, out_dtype=torch.float32) - need_logical_ids = get_moe_capture_callback(infer_state, self.layer_num_) is not None - weights, indices, logical_topk_ids = self._select_experts( - logits, infer_state, layer_weight, return_logical_ids=need_logical_ids - ) + weights, indices = self._select_experts(logits, infer_state, layer_weight) # shared expert 必须先于 routed 计算: fp8 路径 (FuseMoeTriton) 的 fused_experts # 是 inplace 的,_routed_experts 返回后 x 已被覆盖为 routed 输出。 # DS4 shared experts also use the config swiglu_limit clamp, matching SGLang's # DeepseekV2MLP(..., swiglu_limit=config.swiglu_limit) path. shared = self._ffn_tp(input=x, infer_state=infer_state, layer_weight=layer_weight) - routed = self._routed_experts(x, weights, indices, logical_topk_ids, infer_state, layer_weight) + routed = self._routed_experts(x, weights, indices, infer_state, layer_weight) + if self.enable_ep_moe: + return routed + shared out = routed + shared - return out if self.enable_ep_moe else self._tpsp_reduce(input=out, infer_state=infer_state) + return self._tpsp_reduce(input=out, infer_state=infer_state) def _select_experts( - self, - logits, - infer_state: DeepseekV4InferStateInfo, - layer_weight: DeepseekV4TransformerLayerWeight, - return_logical_ids: bool = False, + self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight ): M = logits.shape[0] bias = None @@ -610,32 +594,6 @@ def _select_experts( if input_tokens is None: input_tokens = infer_state.input_ids - eplb = None - if layer_weight.experts_.expert_parallel_state is not None: - eplb = layer_weight.experts_.expert_parallel_state.eplb - if infer_state.is_prefill is True and eplb is not None: - if bias_vl is not None: - raise RuntimeError("DeepSeek-V4 EPLB does not support vision routing yet") - from lightllm.models.deepseek_v4.triton_kernel.moe_topk import ( - deepseek_v4_eplb_topk, - ) - - return deepseek_v4_eplb_topk( - logits=logits, - bias=bias, - input_tokens=input_tokens, - hash_indices_table=hash_indices_table, - topk=self.num_experts_per_tok, - routed_scaling_factor=self.routed_scaling_factor, - logical_to_physical_map=eplb.logical_to_physical_map, - logical_replica_count=eplb.logical_replica_count, - expert_counter=eplb.route_counter, - sample_index=eplb.next_sample_index(), - record_load=eplb.recording, - alloc_tensor_func=self.alloc_tensor, - return_logical_ids=return_logical_ids, - ) - weights = self.alloc_tensor((M, self.num_experts_per_tok), dtype=torch.float32, device=logits.device) indices = self.alloc_tensor((M, self.num_experts_per_tok), dtype=indices_dtype, device=logits.device) topk_softplus_sqrt( @@ -649,7 +607,7 @@ def _select_experts( bias_vl, image_token_start, ) - return weights, indices, indices if return_logical_ids else None + return weights, indices class CompressorInfer(BaseLayerInfer): diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 04867ba586..377b1d611a 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -50,10 +50,6 @@ ) from lightllm.utils.log_utils import init_logger from lightllm.distributed.communication_op import dist_group_manager -from lightllm.common.eplb_utils import ( - EPLB_MAX_STAGING_DEPTH, - extract_eplb_expert_tensors, -) logger = init_logger(__name__) @@ -104,7 +100,6 @@ def _get_compress_rates(self, layer_num): def _init_mem_manager(self): layer_num = self.config["n_layer"] + get_added_mtp_kv_layer_num() state_mtp_step = 0 if self.args.run_mode == "prefill" else self.args.mtp_step - reservations = self._get_post_profile_memory_reservations() retain = self.config["sliding_window"] + DSV4_PROMPT_CACHE_PAGE_SIZE decode_width = max(state_mtp_step + 1, 2 * state_mtp_step) # The retained window may straddle one extra physical page. @@ -132,47 +127,10 @@ def _init_mem_manager(self): else self.args.cpu_cache_token_page_size ), mem_fraction=self.mem_fraction, - memory_reservations=reservations, ) self.req_manager.bind_mem_manager(self.mem_manager) return - def _get_post_profile_memory_reservations(self): - """Only buffers which are not yet visible to cuda.mem_get_info().""" - weights = self._get_eplb_weights() - staging = _get_eplb_staging_nbytes(weights) - sampling = _get_eplb_sampling_peak_nbytes(weights) - return {name: value for name, value in (("eplb_staging", staging), ("eplb_sampling", sampling)) if value} - - def _get_eplb_weights(self): - if self.is_mtp_draft_model or not self.args.enable_prefill_eplb: - return [] - weights = [] - seen = set() - for layer_weight in self.trans_layers_weight: - experts = getattr(layer_weight, "experts_", None) - state = getattr(experts, "expert_parallel_state", None) - if getattr(state, "eplb", None) is None or id(experts) in seen: - continue - seen.add(id(experts)) - weights.append(experts) - return weights - - def get_mtp_profile_weight_exclusion(self): - """Rows present only in target EPLB; the DSpark draft disables EPLB.""" - total = 0 - seen = set() - for experts in self._get_eplb_weights(): - eplb = experts.expert_parallel_state.eplb - redundant = eplb.num_redundant_experts_per_rank - for _, tensor in extract_eplb_expert_tensors(experts): - key = (tensor.data_ptr(), tensor.numel(), tensor.element_size()) - if key in seen: - continue - seen.add(key) - total += redundant * tensor[0].numel() * tensor.element_size() - return total - def _init_att_backend(self): args = get_env_start_args() if args.llm_kv_type == "None": @@ -587,46 +545,3 @@ def apply_chat_template( if tokenize: return self.tokenizer.encode(prompt, add_special_tokens=False) return prompt - - -def _get_eplb_staging_nbytes(weights) -> int: - """Owned bytes for NIXL's reusable staging rows, excluding live expert weights.""" - if not weights: - return 0 - depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) - redundant = weights[0].expert_parallel_state.eplb.num_redundant_experts_per_rank - one_row_nbytes = sum( - tensor[0].numel() * tensor.element_size() for _, tensor in extract_eplb_expert_tensors(weights[0]) - ) - return depth * redundant * one_row_nbytes - - -def _get_eplb_sampling_peak_nbytes(weights) -> int: - """Peak temporary bytes of EPLB sample collection, excluding route counters. - - _collect_local_samples keeps stack(counters), index_select output and index - temporaries live together. Counters are persistent state and intentionally - excluded. CUDA allocator rounding is represented by 512-byte alignment. - """ - counters = [] - seen = set() - for weight in weights: - state = getattr(weight, "expert_parallel_state", None) - eplb = getattr(state, "eplb", None) - counter = getattr(eplb, "route_counter", None) - if counter is None or id(counter) in seen: - continue - seen.add(id(counter)) - counters.append(counter) - if not counters: - return 0 - first = counters[0] - if any(tuple(counter.shape) != tuple(first.shape) or counter.dtype != first.dtype for counter in counters): - raise ValueError("EPLB route counters must have identical shape and dtype") - align = lambda value: (value + 511) // 512 * 512 - stack_bytes = align(len(counters) * first.numel() * first.element_size()) - # index_select has the stack shape. The int64 selection index is bounded - # by one ring axis; this is the selection peak, not a claim that every - # arithmetic intermediate remains live. - index_temp_bytes = align(first.shape[0] * 8) - return 2 * stack_bytes + index_temp_bytes diff --git a/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu deleted file mode 100644 index dc9b53a4b2..0000000000 --- a/lightllm/models/deepseek_v4/triton_kernel/csrc/moe_topk_eplb.cu +++ /dev/null @@ -1,219 +0,0 @@ -// Copyright 2026 LightLLM Team -// SPDX-License-Identifier: Apache-2.0 -// -// DeepSeek-V4 EPLB top-k. The warp-per-row selection follows the public -// Apache-2.0 vLLM topk_softplus_sqrt CUDA implementation, specialized to the -// DSV4 fp32, 256-expert route. - -#include -#include -#include -#include -#include -#include -#include - -namespace { -constexpr int kExperts = 256; -constexpr int kWarps = 4; -constexpr unsigned kMask = 0xffffffffu; - -__device__ __forceinline__ float score(float x) { - return sqrtf(fmaxf(x, 0.f) + log1pf(expf(-fabsf(x)))); -} - -__device__ __forceinline__ bool better(float value, int id, float other_value, int other_id) { - return value > other_value || (value == other_value && id < other_id); -} - -template -__global__ void moe_topk_eplb_kernel( - const float* __restrict__ logits, const float* __restrict__ bias, - const int64_t* __restrict__ input_ids, const int64_t* __restrict__ tid2eid, - float* __restrict__ weights, int64_t* __restrict__ physical_ids, - int64_t* __restrict__ logical_ids, const int32_t* __restrict__ logical_to_physical, - const int32_t* __restrict__ replica_count, int64_t* __restrict__ counter, - int64_t sample_index, int tokens, int topk, int map_slots, int counter_experts, - float routed_scaling_factor) { - const int lane = threadIdx.x & 31; - const int token = blockIdx.x * kWarps + threadIdx.x / 32; - if (token >= tokens) return; - const float* row = logits + token * kExperts; - int selected[8]; - float selected_score[8]; - #pragma unroll - for (int i = 0; i < 8; ++i) { - selected[i] = -1; - selected_score[i] = 0.f; - } - - if constexpr (IsHash) { - if (lane == 0) { - const int64_t input = input_ids[token]; - #pragma unroll - for (int k = 0; k < 8; ++k) { - if (k >= topk) break; - const int id = static_cast(tid2eid[input * topk + k]); - selected[k] = id; - selected_score[k] = score(row[id]); - } - } - } else { - float local_choice[8]; - #pragma unroll - for (int j = 0; j < 8; ++j) { - const int id = lane + j * 32; - local_choice[j] = score(row[id]) + bias[id]; - } - #pragma unroll - for (int k = 0; k < 8; ++k) { - if (k >= topk) break; - float best = -INFINITY; - int best_id = kExperts; - #pragma unroll - for (int j = 0; j < 8; ++j) { - const int id = lane + j * 32; - bool used = false; - #pragma unroll - for (int prev = 0; prev < 8; ++prev) used |= prev < k && selected[prev] == id; - if (!used && better(local_choice[j], id, best, best_id)) { best = local_choice[j]; best_id = id; } - } - for (int offset = 16; offset > 0; offset >>= 1) { - const float other = __shfl_down_sync(kMask, best, offset); - const int other_id = __shfl_down_sync(kMask, best_id, offset); - if (lane + offset < 32 && better(other, other_id, best, best_id)) { best = other; best_id = other_id; } - } - best_id = __shfl_sync(kMask, best_id, 0); - selected[k] = best_id; - selected_score[k] = score(row[best_id]); - } - } - if (lane == 0) { - float sum = 0.f; - #pragma unroll - for (int k = 0; k < 8; ++k) { - if (k < topk) sum += selected_score[k]; - } - #pragma unroll - for (int k = 0; k < 8; ++k) { - if (k >= topk) break; - const int id = selected[k]; - uint32_t replica = 0; - if constexpr (!SingleToken) { - const uint32_t token_hash = static_cast(token) * 2654435769u; - const uint32_t expert_hash = static_cast(id) * 2246822519u; - replica = (token_hash + expert_hash) % static_cast(replica_count[id]); - } - weights[token * topk + k] = selected_score[k] / fmaxf(sum, 1e-20f) * routed_scaling_factor; - physical_ids[token * topk + k] = logical_to_physical[id * map_slots + replica]; - if constexpr (ReturnLogical) logical_ids[token * topk + k] = id; - if constexpr (RecordLoad) { - atomicAdd( - reinterpret_cast(counter + sample_index * counter_experts + id), - 1ULL); - } - } - } -} - -template -void launch(const torch::Tensor& logits, const torch::Tensor& bias, const torch::Tensor& input_ids, - const torch::Tensor& tid2eid, const torch::Tensor& weights, const torch::Tensor& physical_ids, - const torch::Tensor& logical_ids, const torch::Tensor& logical_to_physical, - const torch::Tensor& replica_count, const torch::Tensor& counter, int64_t sample_index, - float routed_scaling_factor) { - const int tokens = logits.size(0), topk = weights.size(1); - moe_topk_eplb_kernel - <<< (tokens + kWarps - 1) / kWarps, kWarps * 32, 0, at::cuda::getCurrentCUDAStream() >>>( - logits.data_ptr(), bias.data_ptr(), input_ids.data_ptr(), tid2eid.data_ptr(), - weights.data_ptr(), physical_ids.data_ptr(), logical_ids.data_ptr(), - logical_to_physical.data_ptr(), replica_count.data_ptr(), counter.data_ptr(), - sample_index, tokens, topk, logical_to_physical.size(1), counter.size(1), routed_scaling_factor); - C10_CUDA_KERNEL_LAUNCH_CHECK(); -} - -template -void dispatch(const torch::Tensor& logits, const torch::Tensor& bias, const torch::Tensor& input_ids, - const torch::Tensor& tid2eid, const torch::Tensor& weights, const torch::Tensor& physical_ids, - const torch::Tensor& logical_ids, const torch::Tensor& logical_to_physical, - const torch::Tensor& replica_count, const torch::Tensor& counter, int64_t sample_index, - bool is_hash, bool single_token, float routed_scaling_factor) { - if (is_hash && single_token) { - launch(logits, bias, input_ids, tid2eid, weights, physical_ids, - logical_ids, logical_to_physical, replica_count, counter, - sample_index, routed_scaling_factor); - } else if (is_hash) { - launch(logits, bias, input_ids, tid2eid, weights, physical_ids, - logical_ids, logical_to_physical, replica_count, counter, - sample_index, routed_scaling_factor); - } else if (single_token) { - launch(logits, bias, input_ids, tid2eid, weights, physical_ids, - logical_ids, logical_to_physical, replica_count, counter, - sample_index, routed_scaling_factor); - } else { - launch(logits, bias, input_ids, tid2eid, weights, physical_ids, - logical_ids, logical_to_physical, replica_count, counter, - sample_index, routed_scaling_factor); - } -} -} // namespace - -void moe_topk_eplb(torch::Tensor logits, torch::Tensor bias, torch::Tensor input_ids, torch::Tensor tid2eid, - torch::Tensor weights, torch::Tensor physical_ids, torch::Tensor logical_ids, - torch::Tensor logical_to_physical, torch::Tensor replica_count, torch::Tensor counter, - int64_t sample_index, bool record_load, bool return_logical_ids, bool is_hash, - float routed_scaling_factor) { - const auto check = [&](const torch::Tensor& tensor, const char* name, torch::ScalarType dtype) { - TORCH_CHECK(tensor.is_cuda(), name, " must be CUDA"); - TORCH_CHECK(tensor.device() == logits.device(), name, " must share logits device"); - TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous"); - TORCH_CHECK(tensor.scalar_type() == dtype, name, " has unexpected dtype"); - }; - check(logits, "logits", torch::kFloat); - check(bias, "bias", torch::kFloat); - check(input_ids, "input_ids", torch::kLong); - check(tid2eid, "tid2eid", torch::kLong); - check(weights, "weights", torch::kFloat); - check(physical_ids, "physical_ids", torch::kLong); - check(logical_ids, "logical_ids", torch::kLong); - check(logical_to_physical, "logical_to_physical", torch::kInt); - check(replica_count, "replica_count", torch::kInt); - check(counter, "counter", torch::kLong); - TORCH_CHECK(logits.dim() == 2 && logits.size(1) == kExperts, "logits must be [M, 256]"); - const int64_t tokens = logits.size(0); - const int64_t topk = weights.size(1); - TORCH_CHECK(topk >= 1 && topk <= 8, "topk must be in [1, 8]"); - TORCH_CHECK(weights.dim() == 2 && weights.size(0) == tokens, "weights must be [M, K]"); - TORCH_CHECK(physical_ids.sizes() == weights.sizes(), "physical_ids must be [M, K]"); - TORCH_CHECK(logical_ids.sizes() == weights.sizes(), "logical_ids must be [M, K]"); - TORCH_CHECK(logical_to_physical.dim() == 2 && logical_to_physical.size(0) == kExperts && - logical_to_physical.size(1) > 0, - "logical_to_physical must be [256, map_slots]"); - TORCH_CHECK(replica_count.dim() == 1 && replica_count.size(0) == kExperts, "replica_count must be [256]"); - TORCH_CHECK(counter.dim() == 2 && counter.size(0) > 0 && counter.size(1) == kExperts, - "counter must be [rows, 256]"); - TORCH_CHECK(sample_index >= 0 && sample_index < counter.size(0), "sample_index outside counter rows"); - c10::cuda::CUDAGuard guard(logits.device()); - const bool single = tokens == 1; - if (is_hash) { - TORCH_CHECK(input_ids.dim() == 1 && input_ids.size(0) == tokens, "input_ids must be [M]"); - TORCH_CHECK(tid2eid.dim() == 2 && tid2eid.size(1) == topk, "tid2eid must be [vocab, K]"); - } else { - TORCH_CHECK(bias.dim() == 1 && bias.size(0) == kExperts, "bias must be [256]"); - } - if (record_load && return_logical_ids) { - dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, - logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); - } else if (record_load) { - dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, - logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); - } else if (return_logical_ids) { - dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, - logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); - } else { - dispatch(logits, bias, input_ids, tid2eid, weights, physical_ids, logical_ids, - logical_to_physical, replica_count, counter, sample_index, is_hash, single, routed_scaling_factor); - } -} - -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("moe_topk_eplb", &moe_topk_eplb); } diff --git a/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py b/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py deleted file mode 100644 index 800055e9a4..0000000000 --- a/lightllm/models/deepseek_v4/triton_kernel/moe_topk.py +++ /dev/null @@ -1,81 +0,0 @@ -import functools -import hashlib -import os - -import torch - - -@torch.no_grad() -def deepseek_v4_eplb_topk( - logits: torch.Tensor, - bias: torch.Tensor | None, - input_tokens: torch.Tensor | None, - hash_indices_table: torch.Tensor | None, - topk: int, - routed_scaling_factor: float, - logical_to_physical_map: torch.Tensor, - logical_replica_count: torch.Tensor, - expert_counter: torch.Tensor, - sample_index: int, - record_load: bool, - alloc_tensor_func=torch.empty, - return_logical_ids: bool = False, -): - """Select DeepSeek-V4 routes and emit EPLB physical expert IDs.""" - token_num, num_experts = logits.shape - if num_experts != 256: - raise RuntimeError(f"DeepSeek-V4 EPLB fused top-k requires 256 experts, got {num_experts}") - if hash_indices_table is not None: - topk = hash_indices_table.shape[1] - weights = alloc_tensor_func((token_num, topk), dtype=torch.float32, device=logits.device) - physical_ids = alloc_tensor_func((token_num, topk), dtype=torch.long, device=logits.device) - logical_ids = ( - alloc_tensor_func((token_num, topk), dtype=torch.long, device=logits.device) if return_logical_ids else None - ) - if token_num == 0: - return weights, physical_ids, logical_ids - _load_cuda().moe_topk_eplb( - logits.contiguous(), - bias.contiguous() if bias is not None else logits, - input_tokens.contiguous() if input_tokens is not None else physical_ids, - hash_indices_table.contiguous() if hash_indices_table is not None else physical_ids, - weights, - physical_ids, - logical_ids if logical_ids is not None else physical_ids, - logical_to_physical_map, - logical_replica_count, - expert_counter, - sample_index, - record_load, - return_logical_ids, - hash_indices_table is not None, - routed_scaling_factor, - ) - return weights, physical_ids, logical_ids - - -@functools.lru_cache(maxsize=1) -def _load_cuda(): - from torch.utils.cpp_extension import load - - source_path = os.path.join(os.path.dirname(__file__), "csrc", "moe_topk_eplb.cu") - flags = ["-O3"] - with open(source_path, "rb") as source_file: - source = source_file.read() - capability = torch.cuda.get_device_capability() - cache_key = b"\0".join( - [ - source, - " ".join(flags).encode(), - torch.__version__.encode(), - str(torch.version.cuda).encode(), - f"sm{capability[0]}{capability[1]}".encode(), - os.environ.get("TORCH_CUDA_ARCH_LIST", "").encode(), - ] - ) - return load( - name=f"lightllm_dsv4_eplb_topk_v1_{hashlib.sha256(cache_key).hexdigest()[:16]}", - sources=[source_path], - extra_cuda_cflags=flags, - verbose=False, - ) diff --git a/lightllm/models/gemma4/tokenizer.py b/lightllm/models/gemma4/tokenizer.py index 760203b8da..5a675856f2 100644 --- a/lightllm/models/gemma4/tokenizer.py +++ b/lightllm/models/gemma4/tokenizer.py @@ -78,9 +78,7 @@ def encode(self, prompt, multimodal_params: MultimodalParams = None, add_special if not input_ids or input_ids[-1] != self.boi_token_index: input_ids.append(self.boi_token_index) img.start_idx = len(input_ids) - img.block_start_idx = img.start_idx input_ids.extend(range(img.token_id, img.token_id + img.token_num)) - img.block_end_idx = len(input_ids) input_ids.append(self.eoi_token_index) if image_end < len(origin_ids) and origin_ids[image_end] == self.eoi_token_index: diff --git a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py index d5347857b5..7311c4d141 100644 --- a/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/qwen3_moe/layer_infer/transformer_layer_infer.py @@ -6,8 +6,9 @@ from lightllm.models.llama.layer_infer.transformer_layer_infer import LlamaTransformerLayerInfer from lightllm.models.llama.infer_struct import LlamaInferStateInfo from lightllm.models.llama.triton_kernel.rotary_emb import rotary_emb_fwd -from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import use_mega_moe -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.triton_ep_impl import FuseMoeTritonEP +from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_fused_moe_ep import ( + use_sm100_mega_moe, +) from lightllm.utils.dist_utils import get_global_world_size from lightllm.utils.envs_utils import get_env_start_args @@ -137,11 +138,7 @@ def overlap_tpsp_token_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if ( - not self.is_moe - or use_mega_moe(layer_weight.experts.quant_method) - or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_token_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) @@ -253,11 +250,7 @@ def overlap_tpsp_context_forward( infer_state1: LlamaInferStateInfo, layer_weight: Qwen3MOETransformerLayerWeight, ): - if ( - not self.is_moe - or use_mega_moe(layer_weight.experts.quant_method) - or isinstance(layer_weight.experts.fuse_moe_impl, FuseMoeTritonEP) - ): + if not self.is_moe or use_sm100_mega_moe(layer_weight.experts.quant_method): return super().overlap_tpsp_context_forward( input_embdings, input_embdings1, infer_state, infer_state1, layer_weight ) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 160a6b61db..4aee85fe6c 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -779,32 +779,15 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: help="""Whether to enable ep moe for deepseekv3 model.""", ) parser.add_argument( - "--ep_moe_backend", + "--ep_redundancy_expert_config_path", type=str, - choices=["auto", "triton"], - default="auto", - help=( - "EP MoE execution backend. 'auto' keeps the existing backend selection; " - "'triton' selects the single-node SM90 FP8 Triton peer backend on Prefill nodes " - "for experts resolved by --expert_dtype fp8." - ), - ) - parser.add_argument( - "--disable_ep_balance_monitor", - action="store_true", - help="""Disable the prefill expert balance monitor enabled by default for EP-MoE.""", + default=None, + help="""Path of the redundant expert config. It can be used for deepseekv3 model.""", ) parser.add_argument( - "--enable_prefill_eplb", + "--auto_update_redundancy_expert", action="store_true", - help="""Enable online expert load balancing for prefill only.""", - ) - parser.add_argument( - "--eplb_num_redundant_experts_per_rank", - type=int, - default=2, - help="""Number of redundant physical experts per EP rank for each MoE layer used by prefill EPLB. - The value must be greater than 0.""", + help="""Whether to update the redundant expert for deepseekv3 model by online expert used counter.""", ) parser.add_argument( "--enable_fused_shared_experts", diff --git a/lightllm/server/api_http.py b/lightllm/server/api_http.py index 3184d39420..9f3f7dd19d 100755 --- a/lightllm/server/api_http.py +++ b/lightllm/server/api_http.py @@ -46,13 +46,7 @@ from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args from lightllm.utils.log_utils import init_logger -from lightllm.utils.error_utils import ( - ClientDisconnected, - GenerationError, - InvalidRequestError, - SERVER_BUSY_MESSAGE, - ServerBusyError, -) +from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError from lightllm.server.metrics.manager import MetricClient from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.shm_port_args import get_shm_port_args @@ -187,11 +181,6 @@ async def invalid_request_exception_handler(request: Request, exc: InvalidReques return create_error_response(HTTPStatus.BAD_REQUEST, str(exc)) -@app.exception_handler(GenerationError) -async def generation_exception_handler(request: Request, exc: GenerationError) -> JSONResponse: - return create_error_response(HTTPStatus.INTERNAL_SERVER_ERROR, str(exc)) - - @app.get("/liveness") @app.post("/liveness") def liveness(): diff --git a/lightllm/server/api_http_pd.py b/lightllm/server/api_http_pd.py index 500ac3df3e..1d8b2112fc 100644 --- a/lightllm/server/api_http_pd.py +++ b/lightllm/server/api_http_pd.py @@ -8,10 +8,10 @@ ``g_objs`` 在 handler 内懒导入,避免与 api_http 循环依赖。 """ +import asyncio import pickle import ujson as json -from anyio import fail_after from fastapi import APIRouter, WebSocket, WebSocketDisconnect from lightllm.server.pd_io_struct import ObjType @@ -33,21 +33,18 @@ async def register_and_keep_alive(websocket: WebSocket): logger.info(f"Client connected from IP: {client_ip}, Port: {client_port}") regist_json = json.loads(await websocket.receive_text()) logger.info(f"received regist_json {regist_json}") - pd_client = await g_objs.httpserver_manager.register_pd(regist_json, websocket) + await g_objs.httpserver_manager.register_pd(regist_json, websocket) try: heartbeat_timeout_seconds = 30 while True: - # Avoid creating a new Task per message so queued PD token packs - # can be drained without extra event-loop scheduling. - with fail_after(heartbeat_timeout_seconds): - data = await websocket.receive_bytes() + data = await asyncio.wait_for(websocket.receive_bytes(), timeout=heartbeat_timeout_seconds) obj = pickle.loads(data) if isinstance(obj, tuple) and obj and obj[0] == ObjType.HEARTBEAT: continue await g_objs.httpserver_manager.put_to_handle_queue(obj) - except TimeoutError: + except asyncio.TimeoutError: logger.warning(f"client {regist_json} heartbeat timed out after {heartbeat_timeout_seconds} seconds") try: await websocket.close(code=1011, reason="PD heartbeat timed out") @@ -60,7 +57,7 @@ async def register_and_keep_alive(websocket: WebSocket): logger.exception(str(e)) finally: logger.error(f"client {regist_json} removed") - await g_objs.httpserver_manager.remove_pd(pd_client) + await g_objs.httpserver_manager.remove_pd(regist_json) return diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index 62b300ea85..c2c2bebdf6 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -30,13 +30,7 @@ from .httpserver_for_pd_master.manager import HttpServerManagerForPDMaster from .api_lightllm import lightllm_get_score from lightllm.utils.envs_utils import get_env_start_args, get_lightllm_websocket_max_message_size -from lightllm.utils.error_utils import ( - ClientDisconnected, - GenerationError, - InvalidRequestError, - SERVER_BUSY_MESSAGE, - ServerBusyError, -) +from lightllm.utils.error_utils import ClientDisconnected, InvalidRequestError, SERVER_BUSY_MESSAGE, ServerBusyError from lightllm.utils.log_utils import init_logger from lightllm.server.metrics.manager import MetricClient @@ -119,35 +113,6 @@ def _serialize_sse_chunk(chunk, choice_nulls=(), response_nulls=()): return json.dumps(d, ensure_ascii=False) -def _serialize_chat_sse_chunk( - request_id, - created, - model, - choice_index, - delta, - choice_nulls=(), - response_nulls=(), - finish_reason=None, -): - """Serialize a chat streaming chunk without constructing Pydantic models.""" - choice = {"index": int(choice_index), "delta": delta} - if finish_reason is not None: - choice["finish_reason"] = finish_reason - for field in choice_nulls: - choice[field] = None - - chunk = { - "id": request_id, - "object": "chat.completion.chunk", - "created": created, - "model": model, - "choices": [choice], - } - for field in response_nulls: - chunk[field] = None - return json.dumps(chunk, ensure_ascii=False) - - def _process_tool_call_id( tool_call_parser, call_item: ToolCallItem, @@ -544,52 +509,11 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req _final_choice_nulls = ("logprobs", "token_ids", "stop_reason") _first_resp_nulls = ("prompt_token_ids",) - def make_chat_sse_chunk(choice_index, delta, choice_nulls=(), response_nulls=(), finish_reason=None): - payload = _serialize_chat_sse_chunk( - chat_completion_id, - created_time, - request.model, - choice_index, - delta, - choice_nulls, - response_nulls, - finish_reason, - ) - return f"data: {payload}\n\n" - - pending_sse_chunks = [] - pending_sse_chars = 0 - - def append_sse_chunk(chunk): - nonlocal pending_sse_chars - pending_sse_chunks.append(chunk) - pending_sse_chars += len(chunk) - - def pending_sse_limit_reached(): - return len(pending_sse_chunks) >= 32 or pending_sse_chars >= 64 * 1024 - - def pop_sse_chunk(): - nonlocal pending_sse_chars - chunk_count = 0 - chunk_chars = 0 - for pending_chunk in pending_sse_chunks: - if chunk_count and (chunk_count >= 32 or chunk_chars + len(pending_chunk) > 64 * 1024): - break - chunk_count += 1 - chunk_chars += len(pending_chunk) - - chunk = "".join(pending_sse_chunks[:chunk_count]) - del pending_sse_chunks[:chunk_count] - pending_sse_chars -= chunk_chars - return chunk - # Streaming case - async def stream_results_inner() -> AsyncGenerator[bytes, None]: + async def stream_results() -> AsyncGenerator[bytes, None]: has_emitted_tool_calls: Dict[int, bool] = collections.defaultdict(bool) has_emitted_first_chunk: Dict[int, bool] = collections.defaultdict(bool) stream_tool_call_ids: Dict[Tuple[int, int], str] = {} - tool_parser = getattr(g_objs.args, "tool_call_parser", None) or "llama3" - history_tool_calls_cnt = _get_history_tool_calls_cnt(request) if tool_parser == "kimi_k2" else 0 from .req_id_generator import convert_sub_id_to_group_id prompt_tokens = 0 @@ -601,25 +525,26 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: completion_tokens += 1 group_request_id = convert_sub_id_to_group_id(sub_req_id) choice_index = sub_req_id - group_request_id - pd_stream_batch_end = metadata.pop("_pd_stream_batch_end", True) delta = request_output current_finish_reason = finish_status.get_finish_reason() - if current_finish_reason == "error" and completion_tokens == 1: - raise GenerationError("Generation failed before producing output") # Emit the initial role-only chunk once per choice, as required by the # OpenAI SSE spec: role appears only in the first delta with content="". if not has_emitted_first_chunk[choice_index]: has_emitted_first_chunk[choice_index] = True - append_sse_chunk( - make_chat_sse_chunk( - choice_index, - {"role": "assistant", "content": ""}, - _first_choice_nulls, - _first_resp_nulls, - ) + first_choice = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(role="assistant", content=""), + finish_reason=None, ) + first_chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[first_choice], + ) + yield f"data: {_serialize_sse_chunk(first_chunk, _first_choice_nulls, _first_resp_nulls)}\n\n" # Handle reasoning content if get_env_start_args().reasoning_parser: @@ -628,9 +553,18 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: ) if reasoning_text: if request.separate_reasoning: - append_sse_chunk( - make_chat_sse_chunk(choice_index, {"reasoning": reasoning_text}, _choice_nulls) + choice_data = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(reasoning=reasoning_text), + finish_reason=None, + ) + chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + choices=[choice_data], + model=request.model, ) + yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" else: delta = reasoning_text + (delta or "") @@ -642,11 +576,22 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: # 1) if there's normal_text, output it as normal content if normal_text and (normal_text.strip() or not has_emitted_tool_calls[sub_req_id]): - append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": normal_text}, _choice_nulls)) + choice_data = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(content=normal_text), + finish_reason=None, + ) + chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + choices=[choice_data], + model=request.model, + ) + yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" # 2) if we found calls, we output them as separate chunk(s) - if calls: - fc_parser = parser_dict[choice_index] + history_tool_calls_cnt = _get_history_tool_calls_cnt(request) + fc_parser = parser_dict[choice_index] for call_item in calls: has_emitted_tool_calls[sub_req_id] = True # transform call_item -> FunctionResponse + ToolCall @@ -668,6 +613,7 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: remaining_call = expected_call.replace(actual_call, "", 1) call_item.parameters = remaining_call + tool_parser = getattr(g_objs.args, "tool_call_parser", None) or "llama3" stream_index = getattr(call_item, "tool_index", None) id_key = (choice_index, stream_index) if call_item.name: @@ -704,7 +650,7 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: choices=[head_choice], model=request.model, ) - append_sse_chunk(f"data: {_serialize_sse_chunk(head_chunk, _choice_nulls)}\n\n") + yield f"data: {_serialize_sse_chunk(head_chunk, _choice_nulls)}\n\n" for arg_delta in _split_tool_argument_delta(call_item.parameters): arg_tool_call = ToolCall( @@ -722,7 +668,7 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: choices=[arg_choice], model=request.model, ) - append_sse_chunk(f"data: {_serialize_sse_chunk(arg_chunk, _choice_nulls)}\n\n") + yield f"data: {_serialize_sse_chunk(arg_chunk, _choice_nulls)}\n\n" else: tool_call = ToolCall( id=tool_call_id if is_tool_head else None, @@ -748,31 +694,38 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: choices=[choice_data], model=request.model, ) - append_sse_chunk(f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n") + yield f"data: {_serialize_sse_chunk(chunk, _choice_nulls)}\n\n" else: if delta: # If this is the final token, merge content with finish_reason if current_finish_reason is not None: if has_emitted_tool_calls[sub_req_id] and current_finish_reason == "stop": current_finish_reason = "tool_calls" - append_sse_chunk( - make_chat_sse_chunk( - choice_index, - {"content": delta}, - _final_choice_nulls, - finish_reason=current_finish_reason, - ) + delta_message = DeltaMessage(content=delta) + stream_choice = ChatCompletionStreamResponseChoice( + index=choice_index, delta=delta_message, finish_reason=current_finish_reason ) - if pd_stream_batch_end: - while pending_sse_chunks: - yield pop_sse_chunk() - else: - while pending_sse_limit_reached(): - yield pop_sse_chunk() + stream_resp = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[stream_choice], + ) + yield f"data: {_serialize_sse_chunk(stream_resp, _final_choice_nulls)}\n\n" # Skip the separate final-chunk logic below continue else: - append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": delta}, _choice_nulls)) + delta_message = DeltaMessage(content=delta) + stream_choice = ChatCompletionStreamResponseChoice( + index=choice_index, delta=delta_message, finish_reason=None + ) + stream_resp = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[stream_choice], + ) + yield f"data: {_serialize_sse_chunk(stream_resp, _choice_nulls)}\n\n" # Emit a per-choice final chunk with finish_reason (for tool_calls path # or when no delta was emitted alongside finish_reason). @@ -785,33 +738,52 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: flush_reasoning, flush_text = parser.flush() if flush_reasoning: if request.separate_reasoning: - flush_delta = {"reasoning": flush_reasoning} + flush_choice = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(reasoning=flush_reasoning), + finish_reason=None, + ) else: # vLLM compat: emit buffered thinking as content - flush_delta = {"content": flush_reasoning} - append_sse_chunk(make_chat_sse_chunk(choice_index, flush_delta, _choice_nulls)) + flush_choice = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(content=flush_reasoning), + finish_reason=None, + ) + flush_chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[flush_choice], + ) + yield f"data: {_serialize_sse_chunk(flush_chunk, _choice_nulls)}\n\n" if flush_text: - append_sse_chunk(make_chat_sse_chunk(choice_index, {"content": flush_text}, _choice_nulls)) + flush_choice = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(content=flush_text), + finish_reason=None, + ) + flush_chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[flush_choice], + ) + yield f"data: {_serialize_sse_chunk(flush_chunk, _choice_nulls)}\n\n" if has_emitted_tool_calls[sub_req_id] and current_finish_reason == "stop": current_finish_reason = "tool_calls" - append_sse_chunk( - make_chat_sse_chunk( - choice_index, - {}, - _final_choice_nulls, - finish_reason=current_finish_reason, - ) + final_choice = ChatCompletionStreamResponseChoice( + index=choice_index, + delta=DeltaMessage(), + finish_reason=current_finish_reason, ) - - if pd_stream_batch_end: - while pending_sse_chunks: - yield pop_sse_chunk() - else: - while pending_sse_limit_reached(): - yield pop_sse_chunk() - - while pending_sse_chunks: - yield pop_sse_chunk() + final_chunk = ChatCompletionStreamResponse( + id=chat_completion_id, + created=created_time, + model=request.model, + choices=[final_choice], + ) + yield f"data: {_serialize_sse_chunk(final_chunk, _final_choice_nulls)}\n\n" reasoning_parser = get_env_start_args().reasoning_parser completion_tokens_details = None @@ -836,17 +808,6 @@ async def stream_results_inner() -> AsyncGenerator[bytes, None]: yield "data: [DONE]\n\n".encode("utf-8") - async def stream_results() -> AsyncGenerator[bytes, None]: - try: - async for chunk in stream_results_inner(): - yield chunk - except ClientDisconnected: - raise - except Exception: - while pending_sse_chunks: - yield pop_sse_chunk() - raise - background_tasks = BackgroundTasks() return CustomStreamingResponse( _safe_stream_wrapper(stream_results()), media_type="text/event-stream", background=background_tasks diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 2f13d9405d..32d9f92114 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -1,6 +1,5 @@ import multiprocessing as mp import os -import tempfile import uuid import subprocess import math @@ -27,7 +26,6 @@ auto_set_response_parsers, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args -from lightllm.utils.device_utils import is_sm100_gpu, is_sm90_gpu logger = init_logger(__name__) @@ -193,22 +191,6 @@ def _launch_subprocesses(args: StartArgs): if args.enable_dp_prefill_balance: assert args.enable_tpsp_mix_mode and args.dp > 1, "need set --enable_tpsp_mix_mode firstly and --dp > 1" - if args.ep_moe_backend == "triton": - assert args.enable_ep_moe, "--ep_moe_backend triton requires --enable_ep_moe" - assert args.run_mode == "prefill", "--ep_moe_backend triton only supports --run_mode prefill" - assert args.nnodes == 1, "--ep_moe_backend triton only supports a single node" - assert not args.enable_prefill_eplb, "--ep_moe_backend triton does not support --enable_prefill_eplb" - assert is_sm90_gpu(), "--ep_moe_backend triton only supports SM90 GPUs" - - if args.enable_prefill_eplb: - assert args.enable_ep_moe, "--enable_prefill_eplb requires --enable_ep_moe" - assert not args.enable_prefill_cudagraph, "--enable_prefill_eplb does not support --enable_prefill_cudagraph" - # EPLB updates expert weights in place, but SM100 Mega-MoE caches transformed weights by tensor data_ptr. - assert not is_sm100_gpu(), "--enable_prefill_eplb does not support SM100" - assert ( - args.eplb_num_redundant_experts_per_rank > 0 - ), "--eplb_num_redundant_experts_per_rank must be greater than 0" - if args.enable_ep_moe: allowed_ep_prefill_att_backends = {"auto", "fa3", "triton", "flashqla"} for backend in args.llm_prefill_att_backend: @@ -487,15 +469,6 @@ def _launch_subprocesses(args: StartArgs): ], ) - instance_disk_cache_dir = None - if args.enable_cpu_cache and args.enable_disk_cache: - cache_base_dir = args.disk_cache_dir or tempfile.gettempdir() - disk_cache_name = os.getenv("DISK_CACHE_NAME") or f"lightllm_disk_cache_{get_unique_server_name()}" - if disk_cache_name in (".", "..") or os.path.basename(disk_cache_name) != disk_cache_name: - raise ValueError("DISK_CACHE_NAME must be a single directory name") - instance_disk_cache_dir = os.path.join(cache_base_dir, disk_cache_name) - process_manager.register_disk_cache_dir(instance_disk_cache_dir) - if args.enable_cpu_cache: from .multi_level_kv_cache.manager import start_multi_level_kv_cache_manager @@ -503,7 +476,7 @@ def _launch_subprocesses(args: StartArgs): start_funcs=[ start_multi_level_kv_cache_manager, ], - start_args=[(args, instance_disk_cache_dir)], + start_args=[(args,)], ) process_manager.start_submodule_processes( diff --git a/lightllm/server/api_stream_obj.py b/lightllm/server/api_stream_obj.py index 8fdbde1355..e56a24ca17 100644 --- a/lightllm/server/api_stream_obj.py +++ b/lightllm/server/api_stream_obj.py @@ -19,41 +19,12 @@ headers immediately. """ -import time - from fastapi.responses import StreamingResponse from starlette.types import Send from lightllm.utils.envs_utils import get_env_start_args -_pd_send_next_metric_time = 0.0 -_pd_send_max_duration = 0.0 -_pd_send_max_bytes = 0 - - -def _record_pd_send_metrics(duration, body_size): - global _pd_send_next_metric_time - global _pd_send_max_duration - global _pd_send_max_bytes - - _pd_send_max_duration = max(_pd_send_max_duration, duration) - _pd_send_max_bytes = max(_pd_send_max_bytes, body_size) - now = time.monotonic() - if now < _pd_send_next_metric_time: - return - - from lightllm.server.api_http import g_objs - - g_objs.httpserver_manager.metric_client.histogram_observe( - "lightllm_pd_master_http_send_duration", _pd_send_max_duration - ) - g_objs.httpserver_manager.metric_client.gauge_set("lightllm_pd_master_http_send_bytes", _pd_send_max_bytes) - _pd_send_max_duration = 0.0 - _pd_send_max_bytes = 0 - _pd_send_next_metric_time = now + 1.0 - - class CustomStreamingResponse(StreamingResponse): """Send the HTTP status only after the first body chunk is ready. @@ -68,21 +39,16 @@ class CustomStreamingResponse(StreamingResponse): """ async def stream_response(self, send: Send) -> None: - start_args = get_env_start_args() # Some clients and proxies require headers before first-token work is # complete. Let deployments opt out of the delayed response start. - if start_args.disable_delay_response_start: + if get_env_start_args().disable_delay_response_start: await super().stream_response(send) return - record_pd_send_metrics = start_args.run_mode == "pd_master" async def send_chunk(chunk): if not isinstance(chunk, (bytes, memoryview)): chunk = chunk.encode(self.charset) - send_start = time.monotonic() await send({"type": "http.response.body", "body": chunk, "more_body": True}) - if record_pd_send_metrics: - _record_pd_send_metrics(time.monotonic() - send_start, len(chunk)) async def send_response_start(): # Read status and headers at send time. The first body iteration diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index 1ee9994643..a28e51dd9e 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -133,23 +133,6 @@ def _normalize_multimodal_content_types(messages: list) -> None: part["type"] = "audio" -def _redact_message_image_data_in_place(messages: list) -> None: - for message in messages: - content = message.get("content") - if not isinstance(content, list): - continue - - for part in content: - image_url = part.get("image_url") - if not image_url: - continue - - url = image_url["url"] - if url.startswith("data:image"): - media_type, separator, _ = url.partition(",") - image_url["url"] = f"{media_type}{separator}" if separator else "" - - async def build_prompt(request, tools) -> str: # pydantic格式转成dict, 否则,当根据tokenizer_config.json拼template时,Jinja判断无法识别 messages = [m.model_dump(by_alias=True, exclude_none=True) for m in request.messages] @@ -186,12 +169,9 @@ async def build_prompt(request, tools) -> str: try: input_str = tokenizer.apply_chat_template(**kwargs, tokenize=False, add_generation_prompt=True, tools=tools) except Exception as e: - request_dump = request.model_dump(by_alias=True, exclude_none=True) - _redact_message_image_data_in_place(request_dump["messages"]) - _redact_message_image_data_in_place(kwargs["conversation"]) logger.exception( "Failed to build prompt. request=%s tools=%s template_kwargs=%s", - json.dumps(request_dump, ensure_ascii=False, default=str), + json.dumps(request.model_dump(by_alias=True, exclude_none=True), ensure_ascii=False, default=str), json.dumps(tools, ensure_ascii=False, default=str), json.dumps(kwargs, ensure_ascii=False, default=str), ) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index c2f3d5e042..c14589f3e3 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -173,8 +173,7 @@ class GuidedJsonSchema(ctypes.Structure): def initialize(self, constraint: str, tokenizer): constraint_bytes = constraint.encode("utf-8") - if len(constraint_bytes) >= JSON_SCHEMA_MAX_LENGTH: - raise ValueError("Guided json schema is too long.") + assert len(constraint_bytes) < JSON_SCHEMA_MAX_LENGTH, "Guided json schema is too long." ctypes.memmove(self.constraint, constraint_bytes, len(constraint_bytes)) self.length = len(constraint_bytes) diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index c82d46cfbc..9c5260ac59 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -189,10 +189,8 @@ class StartArgs: default="gpu_counter", metadata={"choices": ["cpu_counter", "pin_mem_counter", "gpu_counter"]} ) enable_ep_moe: bool = field(default=False) - ep_moe_backend: str = field(default="auto", metadata={"choices": ["auto", "triton"]}) - disable_ep_balance_monitor: bool = field(default=False) - enable_prefill_eplb: bool = field(default=False) - eplb_num_redundant_experts_per_rank: int = field(default=2) + ep_redundancy_expert_config_path: Optional[str] = field(default=None) + auto_update_redundancy_expert: bool = field(default=False) enable_fused_shared_experts: bool = field(default=False) mtp_mode: Optional[str] = field( default=None, diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index eb217d6b91..f53cef44a1 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1572,10 +1572,8 @@ def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> StreamingParseResult: """Streaming incremental parsing for DSML format tool calls.""" - overlap = max(len(self.invoke_end_token), len(self.param_end_token)) - 1 - new_input_window = self._buffer[-overlap:] + new_text - has_new_invoke_end = self.invoke_end_token in new_input_window - has_new_param_end = self.param_end_token in new_input_window + overlap = len(self.param_end_token) - 1 + has_new_param_end = self.param_end_token in self._buffer[-overlap:] + new_text self._buffer += new_text normal_text_parts = [] calls: List[ToolCallItem] = [] @@ -1625,7 +1623,7 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami if self.eot_token.startswith(current_text): return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - complete_invoke_match = self.invoke_regex.match(current_text) if has_new_invoke_end else None + complete_invoke_match = self.invoke_regex.match(current_text) if complete_invoke_match: func_name = complete_invoke_match.group(1) invoke_body = complete_invoke_match.group(2) @@ -1725,7 +1723,7 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami "name": func_name, "arguments": {}, } - if has_new_param_end: + else: # Stream arguments as complete parameters are parsed param_matches = self.param_regex.findall(partial_body) if param_matches and len(param_matches) > len(self._accumulated_params): diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index 10b1e18ed1..47cfed4c88 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -1,12 +1,10 @@ import asyncio -import time class AsyncQueue: def __init__(self): self.datas = [] self.event = asyncio.Event() - self.oldest_put_time = None async def wait_to_ready(self): try: @@ -18,22 +16,15 @@ async def get_all_data(self): self.event.clear() ans = self.datas self.datas = [] - self.oldest_put_time = None return ans async def put(self, obj): was_empty = not self.datas self.datas.append(obj) if was_empty: - self.oldest_put_time = time.monotonic() self.event.set() return - def oldest_age(self): - if self.oldest_put_time is None: - return 0.0 - return time.monotonic() - self.oldest_put_time - async def wait_to_get_all_data(self): await self.wait_to_ready() handle_list = await self.get_all_data() diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 82c7200286..ec8949d283 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -114,8 +114,6 @@ def __init__( self.tokenizer = get_tokenizer(args.model_dir, args.tokenizer_mode, trust_remote_code=args.trust_remote_code) self.req_id_to_out_inf: Dict[int, ReqStatus] = {} # value type (out_str, metadata, finished, event) - # key 存在表示 PD 请求正在登记,value 表示登记期间是否已收到 ABORT。 - self._pd_registration_abort_flags: Dict[int, bool] = {} self.forwarding_queue: AsyncQueue = None # p d 分离模式使用的转发队列, 需要延迟初始化 self.max_req_total_len = args.max_req_total_len @@ -318,21 +316,6 @@ def alloc_req_id(self, sampling_params): assert False, "dead code path" return group_request_id - def begin_pd_request_registration(self, group_req_id: int) -> None: - self._pd_registration_abort_flags[group_req_id] = False - - def cancel_pd_request_registration(self, group_req_id: int) -> None: - """PD 请求在正式登记前结束时,清理对应的待处理 ABORT。""" - self._pd_registration_abort_flags.pop(group_req_id, None) - - def _register_req_status(self, group_req_id: int, req_status: "ReqStatus") -> None: - """发布 PD 请求,并立即消费登记期间收到的 ABORT。""" - self.req_id_to_out_inf[group_req_id] = req_status - if self._pd_registration_abort_flags.pop(group_req_id, False): - for req in req_status.group_req_objs.shm_req_objs: - req.is_aborted = True - logger.warning(f"applied pending abort for group_request_id {group_req_id}") - async def generate( self, prompt: Union[str, List[int]], @@ -407,7 +390,7 @@ async def generate( ) prompt_tokens = len(prompt_ids) - self._check_and_repair_length(prompt_tokens, sampling_params) + prompt_ids = await self._check_and_repair_length(prompt_ids, sampling_params) # 监控 self.metric_client.counter_inc("lightllm_request_count") self.metric_client.histogram_observe("lightllm_request_input_length", prompt_tokens) @@ -428,7 +411,11 @@ async def generate( await pd_upload_websocket.send( pickle.dumps((ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS, group_request_id, array("q", prompt_ids))) ) - await pd_event.wait() + try: + await asyncio.wait_for(pd_event.wait(), timeout=180) + except asyncio.TimeoutError: + logger.error(f"pd prefill node wait pd_event 180s time out, group_req_id {group_request_id}") + raise Exception(f"group_req_id {group_request_id} wait pd_event time out") decode_node_info: PDDecodeNodeInfo = pd_event.decode_node_info sampling_params.pd_kv_trans_params.set(pickle.dumps(decode_node_info)) @@ -482,7 +469,7 @@ async def generate( ) req_status = ReqStatus(group_request_id, multimodal_params, req_objs, start_time) - self._register_req_status(group_request_id, req_status) + self.req_id_to_out_inf[group_request_id] = req_status # RL:请求已登记到 req_id_to_out_inf 并即将转发下游,从 admission gate # 注销,避免 pause 统计里仍把它算作“等待准入”的 pending 请求。 if self.rl_controller is not None: @@ -697,9 +684,10 @@ def get_real_supported_max_req_total_len(self): self.max_req_total_len, ) - def _check_and_repair_length(self, prompt_tokens: int, sampling_params: SamplingParams) -> None: - if prompt_tokens <= 0: + async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params: SamplingParams): + if not prompt_ids: raise InvalidRequestError("The input prompt must not be empty.") + prompt_tokens = len(prompt_ids) # MTP overlap reserves an additional KV window in get_real_supported_max_req_total_len. real_supported_max_req_total_len = self.get_real_supported_max_req_total_len() @@ -724,7 +712,7 @@ def _check_and_repair_length(self, prompt_tokens: int, sampling_params: Sampling ) # last repaired - req_total_len = prompt_tokens + sampling_params.max_new_tokens + req_total_len = len(prompt_ids) + sampling_params.max_new_tokens if req_total_len > self.max_req_total_len: raise InvalidRequestError( f"This model's maximum context length is {self.max_req_total_len} tokens. " @@ -733,6 +721,8 @@ def _check_and_repair_length(self, prompt_tokens: int, sampling_params: Sampling f"Please reduce the length of the input prompt or the number of requested output tokens." ) + return prompt_ids + async def transfer_to_next_module_or_node( self, prompt: str, @@ -937,10 +927,6 @@ async def _wait_to_token_package( async def abort(self, group_req_id: int) -> bool: req_status: ReqStatus = self.req_id_to_out_inf.get(group_req_id, None) if req_status is None: - if group_req_id in self._pd_registration_abort_flags: - self._pd_registration_abort_flags[group_req_id] = True - logger.warning(f"deferred abort for registering group_request_id {group_req_id}") - return True logger.warning(f"aborted group_request_id {group_req_id} not exist") return False diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 450552a62a..76d61e6583 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -30,9 +30,6 @@ logger = init_logger(__name__) -_PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS = 5 -_PD_RECONNECT_DELAY_SECONDS = 10 - async def timer_log(manager: HttpServerManager): while True: @@ -86,9 +83,10 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O pd_handle_loop 主要负责与 pd master 进行注册连接,然后接收pd master发来的请求,然后 将推理结果转发给 pd master进行处理。 """ + # 创建转发队列 + forwarding_queue = AsyncQueue() + while True: - # 转发队列属于当前连接,避免超时未退出的旧请求在重连后上报过期 token。 - forwarding_queue = AsyncQueue() forwarding_tokens_task = None heartbeat_task = None generation_tasks: Dict[int, asyncio.Task] = {} @@ -132,7 +130,6 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O group_req_id = sampling_params.group_request_id pd_event = asyncio.Event() group_req_id_to_event[group_req_id] = pd_event - manager.begin_pd_request_registration(group_req_id) generation_task = asyncio.create_task( _pd_process_generate( manager=manager, @@ -146,13 +143,9 @@ async def _pd_handle_task(manager: HttpServerManager, pd_master_obj: PD_Master_O ) generation_tasks[group_req_id] = generation_task - def remove_generation_task( - task: asyncio.Task, request_id: int = group_req_id, tasks=generation_tasks - ): - if tasks.get(request_id) is task: - tasks.pop(request_id, None) - # task 可能在首次运行前被取消,此时协程内的 finally 不会执行。 - manager.cancel_pd_request_registration(request_id) + def remove_generation_task(task: asyncio.Task, request_id: int = group_req_id): + if generation_tasks.get(request_id) is task: + generation_tasks.pop(request_id, None) generation_task.add_done_callback(remove_generation_task) elif obj[0] == ObjType.ABORT: @@ -161,7 +154,15 @@ def remove_generation_task( generation_task = generation_tasks.get(group_req_id) if generation_task is not None and not generation_task.done(): generation_task.cancel() - await manager.abort(group_req_id) + if not (await manager.abort(group_req_id)): + + async def delayed_abort_task(group_req_id, retry_count): + for _ in range(retry_count): + await asyncio.sleep(5.0) + if await manager.abort(group_req_id): + break + + asyncio.create_task(delayed_abort_task(group_req_id=group_req_id, retry_count=4)) elif obj[0] == ObjType.PD_REQ_DECODE_NODE_INFO: _, group_req_id, decode_node_info = obj @@ -183,30 +184,15 @@ def remove_generation_task( logger.error("connetion to pd_master has error") logger.exception(str(e)) finally: - # Cancel the connection's requests even if their generators cannot exit promptly. - # abort() also defers cancellation for requests that have not registered shm_req yet. - for group_req_id in generation_tasks: - await manager.abort(group_req_id) child_tasks = [task for task in (forwarding_tokens_task, heartbeat_task) if task is not None] child_tasks.extend(generation_tasks.values()) for task in child_tasks: task.cancel() if child_tasks: - done_tasks, pending_tasks = await asyncio.wait( - child_tasks, timeout=_PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS - ) - if done_tasks: - await asyncio.gather(*done_tasks, return_exceptions=True) - if pending_tasks: - logger.warning( - "timed out after %s seconds cleaning up %s PD child task(s); reconnecting", - _PD_CHILD_TASK_CLEANUP_TIMEOUT_SECONDS, - len(pending_tasks), - ) - for task in pending_tasks: - task.cancel() + await asyncio.gather(*child_tasks, return_exceptions=True) - await asyncio.sleep(_PD_RECONNECT_DELAY_SECONDS) + await asyncio.sleep(10) + await forwarding_queue.get_all_data() logger.info("reconnection to pd_master") @@ -298,30 +284,16 @@ async def _pd_process_generate( ) except Exception: logger.exception(f"report pd node generate error failed, group_request_id: {group_request_id}") - finally: - manager.cancel_pd_request_registration(sampling_params.group_request_id) # 转发token的task async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): max_message_size = get_lightllm_websocket_max_message_size() - event_loop = asyncio.get_running_loop() - next_queue_metric_time = 0.0 while True: - await forwarding_queue.wait_to_ready() - queue_duration = forwarding_queue.oldest_age() - handle_list = await forwarding_queue.get_all_data() + handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: - now = event_loop.time() - if now >= next_queue_metric_time: - from lightllm.server.api_http import g_objs - - g_objs.httpserver_manager.metric_client.histogram_observe( - "lightllm_pd_forward_queue_duration", queue_duration - ) - next_queue_metric_time = now + 1.0 load_info: dict = _get_load_info() pending_handle_lists = [] group_start = 0 diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index 0f08782b45..f00f325ea2 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -60,8 +60,6 @@ def __init__( self.health_timeout = int(os.getenv("HEALTH_TIMEOUT", "200")) self.latest_success_infer_time = time.time() self.running_request_count = 0 - self.next_request_queue_metric_time = 0.0 - self._abort_notify_tasks = set() # 限流开关只在 PD Master 生效;P/D 节点不读取本地开关或超时配置,只执行 Master 下发的值。 self.enable_pd_node_self_request_limit = not args.disable_pd_node_self_request_limit self.pd_node_resource_wait_timeout_seconds = get_pd_node_resource_wait_timeout_seconds() @@ -80,7 +78,7 @@ def __init__( return def get_real_supported_max_req_total_len(self): - # HttpServerManager.generate 会借用长度校验逻辑,其中会调用本方法。 + # HttpServerManager.generate 会借用 _check_and_repair_length(self, ...),其中会调用本方法。 # PD master 无本地 token 池 shm 计数;上限与启动参数及子节点对齐的 max_req_total_len 一致。 return self.max_req_total_len @@ -98,14 +96,11 @@ def is_healthy(self): return False async def register_pd(self, pd_info_json, websocket): - return self.pd_manager.register_pd(pd_info_json, websocket) - - async def remove_pd(self, pd_client: PD_Client_Obj): - self.pd_manager.remove_pd(pd_client) - # Wake every stage so the request's existing error path aborts the surviving peer. - for req_status in self.req_id_to_out_inf.values(): - if req_status.p_node is pd_client or req_status.d_node is pd_client: - await req_status.set_error(f"PD {pd_client.mode} node {pd_client.client_ip_port} disconnected") + self.pd_manager.register_pd(pd_info_json, websocket) + return + + async def remove_pd(self, pd_info_json): + self.pd_manager.remove_pd(pd_info_json) return async def update_req_status(self, upkv_status: PDUpKVStatus): @@ -186,9 +181,12 @@ async def _generate( # 计算输入的 input_token_num, 进行校验,如果输入+输出参数设置太长,则将 # sampling_params 的参数进行修正。 input_token_num = await asyncio.to_thread(self.tokens, prompt, multimodal_params, sampling_params) + fake_prompt_ids = [0 for _ in range(input_token_num)] from lightllm.server.httpserver.manager import HttpServerManager - HttpServerManager._check_and_repair_length(self, prompt_tokens=input_token_num, sampling_params=sampling_params) + await HttpServerManager._check_and_repair_length( + self, prompt_ids=fake_prompt_ids, sampling_params=sampling_params + ) return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) @@ -532,7 +530,7 @@ async def fetch_pd_stream( old_max_new_tokens = sampling_params.max_new_tokens sampling_params.max_new_tokens = 1 - await p_node.send_control_message(pickle.dumps((ObjType.REQ, (prompt, sampling_params, multimodal_params)))) + await p_node.websocket.send_bytes(pickle.dumps((ObjType.REQ, (prompt, sampling_params, multimodal_params)))) try: await self._wait_for_event_or_disconnect( @@ -551,7 +549,7 @@ async def fetch_pd_stream( logger.info(f"group_request_id: {group_request_id} get prefill prompt ids len {len(prompt_ids)}") sampling_params.max_new_tokens = old_max_new_tokens - await d_node.send_control_message( + await d_node.websocket.send_bytes( pickle.dumps((ObjType.REQ, (prompt_ids, sampling_params, MultimodalParams()))) ) @@ -572,7 +570,7 @@ async def fetch_pd_stream( upkv_status: PDUpKVStatus = up_status_event.upkv_status pd_kv_trans_params: bytes = upkv_status.pd_kv_trans_params decode_node_info: PDDecodeNodeInfo = pickle.loads(pd_kv_trans_params) - await p_node.send_control_message( + await p_node.websocket.send_bytes( pickle.dumps((ObjType.PD_REQ_DECODE_NODE_INFO, group_request_id, decode_node_info)) ) @@ -598,41 +596,24 @@ async def fetch_pd_stream( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) - request_queue_duration = req_status.oldest_age() token_list = req_status.pop_all_tokens() - if token_list and now >= self.next_request_queue_metric_time: - self.metric_client.histogram_observe( - "lightllm_pd_master_request_queue_duration", request_queue_duration - ) - self.next_request_queue_metric_time = now + 1.0 - ready_token_list = [] for sub_req_id, request_output, metadata, finish_status in token_list: output_index = metadata.get("count_output_tokens") # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 if output_index == 1: - node_run_mode = metadata.pop("node_mode", None) if first_token_gen is False: first_token_gen = True + node_run_mode = metadata.pop("node_mode", None) if node_run_mode == "prefill": if old_max_new_tokens != 1 and finish_status.is_finished_length(): finish_status = FinishStatus(FinishStatus.NO_FINISH) metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) - elif finish_status.status in ( - FinishStatus.FINISHED_ABORTED, - FinishStatus.FINISHED_ERROR, - ): - metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) + yield sub_req_id, request_output, metadata, finish_status else: continue else: metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - ready_token_list.append((sub_req_id, request_output, metadata, finish_status)) - - for index, token_info in enumerate(ready_token_list): - token_info[2]["_pd_stream_batch_end"] = index == len(ready_token_list) - 1 - yield token_info + yield sub_req_id, request_output, metadata, finish_status return @@ -668,9 +649,6 @@ async def _wait_for_prefill_token_if_needed( prompt_cache_len = metadata.get("prompt_cache_len", 0) req_status.put_tokens_to_front(new_tokens) return prompt_cache_len - if token[3].is_finished(): - req_status.put_tokens_to_front(new_tokens) - return ready_kv_len async def _wait_to_token_package( self, @@ -758,24 +736,15 @@ async def abort( except: pass - async def notify_node(node: Optional[PD_Client_Obj]): - if node is None: - return - try: - await node.send_control_message(pickle.dumps((ObjType.ABORT, group_request_id))) - except BaseException: - pass - - # HTTP request cancellation must not cancel an ABORT while it is waiting for - # the node's websocket send lock. Keep independent tasks alive until both - # nodes have received the cleanup message or their connections fail. - notify_tasks = [asyncio.create_task(notify_node(node)) for node in (p_node, d_node) if node is not None] - self._abort_notify_tasks.update(notify_tasks) - for task in notify_tasks: - task.add_done_callback(self._abort_notify_tasks.discard) + try: + await p_node.websocket.send_bytes(pickle.dumps((ObjType.ABORT, group_request_id))) + except: + pass - if notify_tasks: - await asyncio.gather(*(asyncio.shield(task) for task in notify_tasks)) + try: + await d_node.websocket.send_bytes(pickle.dumps((ObjType.ABORT, group_request_id))) + except: + pass return @@ -797,8 +766,6 @@ async def put_to_handle_queue(self, obj): async def handle_loop(self): self.infos_queues = AsyncQueue() asyncio.create_task(self.timer_log()) - event_loop = asyncio.get_running_loop() - next_ingress_queue_metric_time = 0.0 use_config_server = self.args.config_server_host and self.args.config_server_port @@ -808,15 +775,7 @@ async def handle_loop(self): asyncio.create_task(register_loop(self)) while True: - await self.infos_queues.wait_to_ready() - ingress_queue_duration = self.infos_queues.oldest_age() - objs = await self.infos_queues.get_all_data() - now = event_loop.time() - if objs and now >= next_ingress_queue_metric_time: - self.metric_client.histogram_observe( - "lightllm_pd_master_ingress_queue_duration", ingress_queue_duration - ) - next_ingress_queue_metric_time = now + 1.0 + objs = await self.infos_queues.wait_to_get_all_data() try: for obj in objs: @@ -883,7 +842,6 @@ def __init__(self, req_id, p_node, d_node) -> None: self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() self.out_token_info_list: List[Tuple[int, str, dict, FinishStatus]] = [] - self.oldest_token_time = None self.p_node: PD_Client_Obj = p_node self.d_node: PD_Client_Obj = d_node self.error_info: Optional[str] = None @@ -919,29 +877,19 @@ def append_token(self, token_info: Tuple[int, str, dict, FinishStatus]): was_empty = not self.out_token_info_list self.out_token_info_list.append(token_info) if was_empty: - self.oldest_token_time = time.monotonic() self.event.set() - def oldest_age(self): - if self.oldest_token_time is None: - return 0.0 - return time.monotonic() - self.oldest_token_time - def pop_all_tokens(self): self.event.clear() ans = self.out_token_info_list self.out_token_info_list = [] - self.oldest_token_time = None return ans def put_tokens_to_front(self, token_list: List[Tuple[int, str, dict, FinishStatus]]): if not token_list: return - was_empty = not self.out_token_info_list self.out_token_info_list = token_list + self.out_token_info_list - if was_empty: - self.oldest_token_time = time.monotonic() self.event.set() @@ -1042,15 +990,12 @@ def register_pd(self, pd_info_json, websocket): self.selector.update_nodes(self.prefill_nodes, self.decode_nodes) logger.info(f"mode: {pd_client.mode} url: {pd_client.client_ip_port} registed") - return pd_client + return - def remove_pd(self, pd_client: PD_Client_Obj): - # A closing connection must not remove a newer registration at the same address. - pd_client.websocket = None - if self.url_to_pd_nodes.get(pd_client.client_ip_port) is not pd_client: - return + def remove_pd(self, pd_info_json): + pd_client = PD_Client_Obj(**pd_info_json) - self.url_to_pd_nodes.pop(pd_client.client_ip_port) + self.url_to_pd_nodes.pop(pd_client.client_ip_port, None) self.prefill_nodes = [e for e in self.prefill_nodes if e.client_ip_port != pd_client.client_ip_port] self.decode_nodes = [e for e in self.decode_nodes if e.client_ip_port != pd_client.client_ip_port] @@ -1081,10 +1026,4 @@ def update_node_load_info(self, load_info: Optional[dict]): def select_p_d_node( self, prompt: Union[str, List[int]], sampling_params: SamplingParams, multimodal_params: MultimodalParams ) -> Tuple[PD_Client_Obj, PD_Client_Obj, PDSelectionExtraInfo]: - if not self.prefill_nodes or not self.decode_nodes: - raise ServerBusyError( - "PD nodes unavailable: " - f"registered_prefill={len(self.prefill_nodes)}, registered_decode={len(self.decode_nodes)}", - status_code=503, - ) return self.selector.select_p_d_node(prompt, sampling_params, multimodal_params) diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index 46d92ca5f1..b8f8d6a1f2 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -54,6 +54,8 @@ class CacheAwareConfig: evict_node_batch: int = 10_000 # 每隔 sample_stride 个字符抽 1 个作为前缀树 key,降低匹配开销与内存。 sample_stride: int = 512 + # 初始化前缀树时通过 sys.setrecursionlimit 调大 Python 调用栈深度。 + recursion_limit: int = 4000 class BalanceRelThresholdController: @@ -107,6 +109,7 @@ def __init__(self, config: Optional[CacheAwareConfig] = None) -> None: sample_stride=self.config.sample_stride, max_node_count=self.config.max_node_count, evict_node_batch=self.config.evict_node_batch, + recursion_limit=self.config.recursion_limit, ) self.balance_rel_threshold_controller = BalanceRelThresholdController() diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py b/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py index 5833ede354..859b7a7329 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/prompt_cache_tree.py @@ -1,5 +1,6 @@ from __future__ import annotations +import sys import time from dataclasses import dataclass, field from threading import Lock, RLock @@ -38,7 +39,8 @@ class PromptCacheTree: """ 用于 cache-aware 选点的 prompt 前缀缓存树。 - prompt 先按 sample_stride 抽稀成 key,再按单字符建树/匹配。 + prompt 先按 sample_stride 抽稀成 key,再递归按单字符建树/匹配。 + recursion_limit 在初始化时通过 sys.setrecursionlimit 调大 Python 调用栈深度。 整棵树节点数有上限;超限时按 LRU 从叶节点批量删除。 """ @@ -47,12 +49,14 @@ def __init__( sample_stride: int = 512, max_node_count: int = 1_000_000, evict_node_batch: int = 10_000, + recursion_limit: int = 4000, ) -> None: """ Args: sample_stride: 每隔多少个字符抽 1 个作为 trie key。 max_node_count: 树中允许的最大节点数(不含 root);超限时触发 LRU 驱逐。 evict_node_batch: 每次驱逐时在超出量基础上额外腾出的节点缓冲数。 + recursion_limit: 初始化时通过 sys.setrecursionlimit 设置的调用栈深度上限。 """ if sample_stride < 1: raise ValueError(f"sample_stride must be >= 1, got {sample_stride}") @@ -60,9 +64,14 @@ def __init__( raise ValueError(f"max_node_count must be >= 0, got {max_node_count}") if evict_node_batch < 1: raise ValueError(f"evict_node_batch must be >= 1, got {evict_node_batch}") + if recursion_limit < 1: + raise ValueError(f"recursion_limit must be >= 1, got {recursion_limit}") self.sample_stride = sample_stride self.max_node_count = max_node_count self.evict_node_batch = evict_node_batch + self.recursion_limit = recursion_limit + if recursion_limit > sys.getrecursionlimit(): + sys.setrecursionlimit(recursion_limit) self.root = _PromptCacheNode() self._node_count = 0 self._leaf_lru: SortedDict[int, _PromptCacheNode] = SortedDict() @@ -96,37 +105,32 @@ def insert(self, text: str, prefill_node: str) -> None: self._evict_if_needed() def _insert_at(self, node: _PromptCacheNode, key: str, depth: int, prefill_node: str) -> None: - path = [] try: - while True: - path.append(node) - if depth >= len(key): - break - - ch = key[depth] - child = node.children.get(ch) - if child is None: - child = _PromptCacheNode( - parent=node, - edge_char=ch, - last_insert_time=time.monotonic(), - ) - child.last_time_mark = self._gen_time_mark() - node.children[ch] = child - self._node_count += 1 - - node = child - depth += 1 + if depth >= len(key): + return + + ch = key[depth] + child = node.children.get(ch) + if child is None: + child = _PromptCacheNode( + parent=node, + edge_char=ch, + last_insert_time=time.monotonic(), + ) + child.last_time_mark = self._gen_time_mark() + node.children[ch] = child + self._node_count += 1 + + self._insert_at(child, key, depth + 1, prefill_node) finally: - for path_node in reversed(path): - if path_node is not self.root: - path_node.last_prefill_node = prefill_node - path_node.last_insert_time = time.monotonic() - if path_node.last_time_mark in self._leaf_lru: - self._leaf_lru.pop(path_node.last_time_mark, None) - path_node.last_time_mark = self._gen_time_mark() - if self._is_leaf(path_node): - self._leaf_lru[path_node.last_time_mark] = path_node + if node is not self.root: + node.last_prefill_node = prefill_node + node.last_insert_time = time.monotonic() + if node.last_time_mark in self._leaf_lru: + self._leaf_lru.pop(node.last_time_mark, None) + node.last_time_mark = self._gen_time_mark() + if self._is_leaf(node): + self._leaf_lru[node.last_time_mark] = node def _gen_time_mark(self) -> int: with self._time_mark_lock: @@ -174,13 +178,14 @@ def prefix_match(self, text: str) -> PromptCacheMatchResult: ) def _match_at(self, node: _PromptCacheNode, key: str, depth: int) -> Tuple[_PromptCacheNode, int]: - while depth < len(key): - child = node.children.get(key[depth]) - if child is None: - break - node = child - depth += 1 - return node, depth + if depth >= len(key): + return node, depth + + child = node.children.get(key[depth]) + if child is None: + return node, depth + + return self._match_at(child, key, depth + 1) def evict_lru_nodes(self) -> int: """节点数超上限时,按 LRU 从叶节点批量删除。 diff --git a/lightllm/server/metrics/metrics.py b/lightllm/server/metrics/metrics.py index b0482f93f6..0d42462c3f 100644 --- a/lightllm/server/metrics/metrics.py +++ b/lightllm/server/metrics/metrics.py @@ -27,25 +27,11 @@ "lightllm_cache_ratio": "cache length / input_length", "lightllm_batch_current_max_tokens": "dynamic max token used for current batch", "lightllm_request_mtp_avg_token_per_step": "Average number of tokens per step", - "lightllm_pd_forward_queue_duration": "Time tokens wait on a P/D node before websocket forwarding (s)", - "lightllm_pd_master_ingress_queue_duration": "Time token packets wait in the PD master ingress queue (s)", - "lightllm_pd_master_request_queue_duration": "Time tokens wait for their PD master request consumer (s)", - "lightllm_pd_master_http_send_duration": "Maximum ASGI body send duration in the latest sample window (s)", - "lightllm_pd_master_http_send_bytes": "Largest ASGI body sent in the latest sample window (bytes)", "lightllm_prompt_tokens_total": "Total number of prefill tokens processed", "lightllm_generation_tokens_total": "Total number of generation tokens processed", "lightllm_cache_hit_rate": "Prefix cache hit rate of latest completed request", "lightllm_gen_throughput": "Generation throughput of latest completed request (tokens/s)", "lightllm_num_running_reqs": "Number of running requests", - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": ( - "Estimated critical-path excess compute per logical routed source token, GFLOPs/token" - ), - "lightllm_prefill_ep_compute_critical_overhead_ratio": ( - "Estimated excess critical compute divided by balanced compute; 0.3 means +30%" - ), - "lightllm_prefill_ep_placement_pressure_drift": ( - "Normalized temporal drift of overloaded-rank pressure from the latest complete prefill report" - ), } @@ -102,15 +88,10 @@ def init_metrics(self, args): self.create_histogram("lightllm_request_mean_time_per_token_duration", self.duration_buckets) self.create_histogram("lightllm_request_first_token_duration", self.duration_buckets) self.create_histogram("lightllm_request_queue_duration_bucket", self.duration_buckets) - self.create_histogram("lightllm_pd_forward_queue_duration", self.duration_buckets) - self.create_histogram("lightllm_pd_master_ingress_queue_duration", self.duration_buckets) - self.create_histogram("lightllm_pd_master_request_queue_duration", self.duration_buckets) - self.create_histogram("lightllm_pd_master_http_send_duration", self.duration_buckets) self.create_histogram("lightllm_batch_inference_duration_bucket", self.duration_buckets, labelnames=["method"]) self.gateway_url = args.metric_gateway self.create_gauge("lightllm_queue_size") - self.create_gauge("lightllm_pd_master_http_send_bytes") self.create_gauge("lightllm_batch_current_size") self.create_gauge("lightllm_batch_pause_size") self.create_gauge("lightllm_batch_current_max_tokens") @@ -130,9 +111,6 @@ def init_metrics(self, args): self.create_gauge("lightllm_cache_hit_rate") self.create_gauge("lightllm_gen_throughput") self.create_gauge("lightllm_num_running_reqs") - self.create_gauge("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token") - self.create_gauge("lightllm_prefill_ep_compute_critical_overhead_ratio") - self.create_gauge("lightllm_prefill_ep_placement_pressure_drift") def create_histogram(self, name, buckets, labelnames=None): all_labels = ["model_name"] + (labelnames or []) diff --git a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py index aff400753d..07a0be319b 100644 --- a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py +++ b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py @@ -1,10 +1,12 @@ import os +import tempfile import time import math from dataclasses import dataclass -from typing import List +from typing import List, Optional import torch +from lightllm.utils.envs_utils import get_unique_server_name from lightllm.utils.log_utils import init_logger from .cpu_cache_client import CpuKvCacheClient @@ -35,7 +37,7 @@ def __init__( self, disk_cache_storage_size: float, cpu_cache_client: CpuKvCacheClient, - cache_dir: str, + disk_cache_dir: Optional[str] = None, ): self.cpu_cache_client = cpu_cache_client self._pages_all_idle = False @@ -48,6 +50,10 @@ def __init__( # 读写同时进行时,分配8线程用来写,16线程用来读 max_concurrent_write_tasks = 8 + if disk_cache_dir: + cache_dir = os.path.join(disk_cache_dir, f"lightllm_disk_cache_{get_unique_server_name()}") + else: + cache_dir = os.path.join(tempfile.gettempdir(), f"lightllm_disk_cache_{get_unique_server_name()}") os.makedirs(cache_dir, exist_ok=True) cache_file = os.path.join(cache_dir, "cache_file") diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index f2b4a576a7..ef5b7369c9 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -9,7 +9,7 @@ import threading import concurrent.futures import setproctitle -from queue import Empty, Queue +from queue import Queue from typing import List from lightllm.server.core.objs import ShmReqManager, Req, StartArgs from lightllm.server.core.objs.io_objs import GroupReqIndexes @@ -27,7 +27,6 @@ class MultiLevelKVCacheManager: def __init__( self, args: StartArgs, - instance_disk_cache_dir, ): self.args: StartArgs = args ports = get_shm_port_args() @@ -46,9 +45,6 @@ def __init__( # 控制进行 cpu cache 页面匹配的时间,超过时间则不再匹配,直接转发。 self.cpu_cache_time_out = 0.5 self.recv_queue = Queue(maxsize=1024) - # ZeroMQ sockets are not thread-safe. Cache workers enqueue completed - # requests so the recv_loop thread remains the sole socket owner. - self.send_to_router_queue = Queue() self.cpu_cache_thread = threading.Thread(target=self.cpu_cache_hanle_loop, daemon=True) self.cpu_cache_thread.start() @@ -60,7 +56,7 @@ def __init__( self.disk_cache_worker = DiskCacheWorker( disk_cache_storage_size=self.args.disk_cache_storage_size, cpu_cache_client=self.cpu_cache_client, - cache_dir=instance_disk_cache_dir, + disk_cache_dir=self.args.disk_cache_dir, ) self.disk_cache_thread = threading.Thread(target=self.disk_cache_worker.run, daemon=True) self.disk_cache_thread.start() @@ -148,7 +144,7 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes # 超时时,放弃进行 cache page 的匹配。 current_time = time.time() if current_time - start_time >= self.cpu_cache_time_out: - self.send_to_router_queue.put(group_req_indexes) + self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) logger.warning( f"cache matching time out {current_time - start_time}s, " f"group_req_id: {group_req_indexes.group_req_id}" @@ -215,23 +211,14 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes for req in reqs: self.shm_req_manager.put_back_req_obj(req) - self.send_to_router_queue.put(group_req_indexes) + self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) return - def _send_finished_group_reqs(self): - while True: - try: - group_req_indexes = self.send_to_router_queue.get_nowait() - except Empty: - return - self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) - def recv_loop(self): try: recv_max_count = 128 while True: - self._send_finished_group_reqs() recv_objs = [] try: # 一次最多从 zmq 中取 recv_max_count 个请求,防止 zmq 队列中请求数量过多导致阻塞了主循环。 @@ -263,7 +250,7 @@ def recv_loop(self): return -def start_multi_level_kv_cache_manager(args, instance_disk_cache_dir, pipe_writer): +def start_multi_level_kv_cache_manager(args, pipe_writer): # 注册graceful 退出的处理 graceful_registry(inspect.currentframe().f_code.co_name) setproctitle.setproctitle(f"lightllm::{get_unique_server_name()}::multi_level_kv_cache") @@ -272,7 +259,6 @@ def start_multi_level_kv_cache_manager(args, instance_disk_cache_dir, pipe_write try: manager = MultiLevelKVCacheManager( args=args, - instance_disk_cache_dir=instance_disk_cache_dir, ) except Exception as e: logger.exception(f"start multi_level_kv_cache_manager has exception {str(e)}") diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index c709c4d0b3..bf39bcb570 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -1,4 +1,3 @@ -import asyncio import enum import time import copy @@ -162,8 +161,6 @@ class PD_Client_Obj: dispatched_prompt_chars: int = 0 # 当前派发到该节点且尚未产出首 token 的请求数。 dispatched_req_num: int = 0 - _send_lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False, repr=False, compare=False) - _send_task: Optional[asyncio.Task] = field(default=None, init=False, repr=False, compare=False) def __post_init__(self): if self.mode not in ["prefill", "decode"]: @@ -175,46 +172,6 @@ def __post_init__(self): def to_llm_url(self): return f"http://{self.client_ip_port}/pd_generate_stream" - async def send_control_message(self, payload: bytes) -> None: - # A disconnected client may still have an old send holding the lock. Do not - # let cleanup messages wait for that send before noticing the invalidation. - if self.websocket is None: - raise ConnectionError(f"PD control connection unavailable: {self.client_ip_port}") - - # Waiting requests remain cancellable BEFORE they advance the compression dictionary. - await self._send_lock.acquire() - try: - if self.websocket is None: - raise ConnectionError(f"PD control connection unavailable: {self.client_ip_port}") - send_task = asyncio.create_task(self.websocket.send_bytes(payload)) - self._send_task = send_task - except BaseException: - self._send_lock.release() - raise - - def finish_send(task: asyncio.Task): - self._send_task = None - try: - task.result() - except BaseException: - self.websocket = None - logger.exception("PD control send failed: peer=%s", self.client_ip_port) - finally: - self._send_lock.release() - - # The connection owns the task AND the lock until the complete frame is sent. - # Hypercorn compresses before awaiting its TCP send lock. Cancelling that wait - # drops the frame but leaves the deflate dictionary advanced for later messages. - send_task.add_done_callback(finish_send) - try: - await asyncio.shield(send_task) - except asyncio.CancelledError: - logger.warning( - "PD control send caller cancelled; connection-owned send continues: " "peer=%s", - self.client_ip_port, - ) - raise - @dataclass class PD_Master_Obj: 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 2aefe7132d..ac254785ba 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -13,9 +13,6 @@ 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 -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - disable_eplb_model_init, -) from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.basemodel.logprobs_manager import PromptLogprobsCaptureManager from lightllm.common.basemodel.moe_route_info_manager import MoeRouteInfoManager @@ -260,11 +257,6 @@ def init_model(self, kvargs): prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling self.profiler = ProcessProfiler(mode=prof_mode, name=prof_name, use_multi_thread=True) if prof_mode else None - if self.args.enable_prefill_eplb: - from lightllm.server.router.model_infer.mode_backend.eplb_manager import EPLBManager - - self.model.eplb_manager = EPLBManager(self.model) - dist.barrier() # 启动infer_loop_thread, 启动两个线程进行推理,对于具备双batch推理折叠得场景 # 可以降低 cpu overhead,大幅提升gpu得使用率。 @@ -353,9 +345,7 @@ def init_mtp_draft_model(self, main_kvargs: dict): model_cfg=draft_model_cfg, spec_mode=spec_mode, ) - with disable_eplb_model_init(): - draft_model = draft_model_class(draft_model_kvargs) - self.draft_models.append(draft_model) + self.draft_models.append(draft_model_class(draft_model_kvargs)) self.logger.info(f"loaded speculative draft model class {self.draft_models[i].__class__}") return diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py index d157d05431..97515faf0d 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 @@ -60,11 +60,6 @@ def infer_loop(self): event_pack.wait_to_forward() - # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal - # collectives/forward. - if self.model.eplb_manager is not None: - self.model.eplb_manager.poll() - self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( 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 96554ccac7..304eede019 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 @@ -127,11 +127,6 @@ def infer_loop(self): event_pack.wait_to_forward() - # EPLB polling performs collectives and may commit weights, so keep all ranks ordered before normal - # collectives/forward. - if self.model.eplb_manager is not None: - self.model.eplb_manager.poll() - self._try_read_new_reqs() prefill_reqs, decode_reqs = self._get_classed_reqs( diff --git a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py b/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py deleted file mode 100644 index 639457fbf5..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/ep_balance_monitor.py +++ /dev/null @@ -1,351 +0,0 @@ -import threading -from array import array -from typing import Optional, Tuple - -import torch -import torch.distributed as dist - -from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters -from lightllm.distributed.communication_op import dist_group_manager -from lightllm.server.metrics.manager import MetricClient -from lightllm.utils.device_utils import is_sm100_gpu -from lightllm.utils.dist_utils import ( - get_global_rank, - get_global_world_size, -) -from lightllm.utils.log_utils import init_logger -from lightllm.utils.shm_port_args import get_shm_port_args - - -logger = init_logger(__name__) - -EP_BALANCE_PREFILL_ROUNDS_PER_REPORT = 100 -EP_BALANCE_ROUND_BUFFER_CAPACITY = 4096 -EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS = 20 -EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD = 0.10 -ROUTE_LOAD = 0 -COMPUTE_LOAD = 1 -GFLOP = 1_000_000_000 - - -def should_enable_ep_balance_monitor(args) -> bool: - if ( - args.enable_prefill_cudagraph - or is_sm100_gpu() - or getattr(dist_group_manager, "ep_mega_moe_buffer", None) is not None - or getattr(dist_group_manager, "ep_triton_moe_buffer", None) is not None - ): - return False - return args.enable_ep_moe and not args.disable_ep_balance_monitor and args.run_mode != "decode" - - -def calculate_prefill_balance_stats( - round_stats: torch.Tensor, # [num_rounds, num_layers, world_size, 2] (route/compute) - layer_routed_experts: torch.Tensor, # [num_layers] - layer_flops_per_expert_token: torch.Tensor, # [num_layers] - layer_topks: torch.Tensor, # [num_layers] - source_token_replication: int, - report_min_route_samples_per_expert: int = 100, -) -> Optional[dict]: - """Summarize complete-prefill samples from [round, layer, rank, route/compute]. - - MoE layers execute sequentially, and every layer waits for its slowest EP - rank. Preserve the layer dimension until after taking the cross-rank max so - that different slow ranks in different layers cannot cancel each other. - """ - assert source_token_replication > 0 - - layer_route_load = round_stats[:, :, :, ROUTE_LOAD].sum(dim=(0, 2)) - minimum_route_samples = layer_routed_experts * report_min_route_samples_per_expert - if torch.any(layer_route_load < minimum_route_samples): - return None - total_route_load = layer_route_load.sum() - - compute_rank_load = round_stats[:, :, :, COMPUTE_LOAD] - compute_rank_load_float = compute_rank_load.to(torch.float64) - excess_compute_load = compute_rank_load_float.max(dim=2).values - compute_rank_load_float.mean(dim=2) - - # Every MoE expert token executes the two projections packed in w13 plus - # the w2 projection. Weight each padded compute token by the layer's - # actual matrix sizes so the metric remains comparable across models. - excess_compute_flops = (excess_compute_load * layer_flops_per_expert_token.to(torch.float64)).sum() - balanced_compute_flops = ( - compute_rank_load_float.mean(dim=2) * layer_flops_per_expert_token.to(torch.float64) - ).sum() - if balanced_compute_flops == 0: - return None - - # Non-TPSP prefill gathers one route-load copy per TP rank. Divide the - # replica count out so GFLOP/token uses logical source tokens. - source_tokens = total_route_load.to(torch.float64) / ( - layer_topks.to(torch.float64).sum() * source_token_replication - ) - if source_tokens == 0: - return None - - return { - "prefill_rounds": int(compute_rank_load.shape[0]), - "critical_overhead_gflops_per_routed_token": float((excess_compute_flops / source_tokens / GFLOP).item()), - "prefill_ep_compute_critical_overhead_ratio": float((excess_compute_flops / balanced_compute_flops).item()), - } - - -def calculate_prefill_placement_pressure_drift( - round_stats: torch.Tensor, - previous_pressure_signature: Optional[torch.Tensor] = None, - bucket_rounds: int = EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, -) -> Tuple[float, torch.Tensor]: - """Measure how prefill rank-pressure placement changes across time buckets. - - This is rank-0 CPU-only report analysis. The returned final bucket is a - compact signature that allows the next report to include the boundary pair. - """ - if bucket_rounds <= 0: - raise ValueError(f"bucket_rounds must be positive, got {bucket_rounds}") - if round_stats.ndim != 4 or round_stats.shape[-1] != 2: - raise ValueError( - "round_stats must have shape [num_rounds, num_layers, world_size, 2], " f"got {tuple(round_stats.shape)}" - ) - num_rounds, num_layers, world_size, _ = round_stats.shape - if num_rounds <= 0: - raise ValueError("round_stats must contain at least one round") - if num_layers <= 0 or world_size <= 0: - raise ValueError("round_stats must contain at least one layer and rank") - if num_rounds % bucket_rounds != 0: - raise ValueError(f"num_rounds ({num_rounds}) must be divisible by bucket_rounds ({bucket_rounds})") - - num_buckets = num_rounds // bucket_rounds - bucket_rank_load = ( - round_stats[:, :, :, COMPUTE_LOAD] - .to(torch.float64) - .reshape(num_buckets, bucket_rounds, num_layers, world_size) - .sum(dim=1) - ) - mean_rank_load = bucket_rank_load.mean(dim=2, keepdim=True).clamp_min(1) - pressure = torch.relu(bucket_rank_load / mean_rank_load - 1) - - if previous_pressure_signature is not None: - expected_shape = (num_layers, world_size) - if tuple(previous_pressure_signature.shape) != expected_shape: - raise ValueError( - "previous_pressure_signature must have shape " - f"{expected_shape}, got {tuple(previous_pressure_signature.shape)}" - ) - left = torch.cat((previous_pressure_signature.to(torch.float64).unsqueeze(0), pressure[:-1]), dim=0) - right = pressure - else: - left = pressure[:-1] - right = pressure[1:] - - total_pressure = (left + right).sum() - if total_pressure == 0: - drift = 0.0 - else: - drift = float((left - right).abs().sum().div(total_pressure).item()) - return drift, pressure[-1].clone() - - -def classify_prefill_placement_pressure_drift(drift: float) -> str: - if drift < EP_BALANCE_PRESSURE_DRIFT_STABLE_THRESHOLD: - return "stable" - return "dynamic" - - -def _find_fused_moe_weights(model): - weights_by_id = {} - for layer in model.trans_layers_weight: - for value in getattr(layer, "__dict__", {}).values(): - if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: - weights_by_id[id(value)] = value - return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) - - -class EPBalanceMonitor: - """Report cross-rank imbalance for non-overlapping blocks of complete prefill rounds.""" - - def __init__(self, model: TpPartBaseModel): - self.global_rank = get_global_rank() - self.world_size = get_global_world_size() - self.weights = _find_fused_moe_weights(model) - self.enabled = bool(self.weights) - - if not self.enabled: - return - - self.source_token_replication = 1 if model.args.enable_tpsp_mix_mode else model.tp_world_size_ - self.counters: list[PrefillEPBalanceCounters] = [PrefillEPBalanceCounters() for _ in self.weights] - for weight, counter in zip(self.weights, self.counters): - weight.fuse_moe_impl.ep_balance_counters = counter - self.layer_routed_experts = torch.tensor( - [weight.n_routed_experts for weight in self.weights], dtype=torch.int64 - ) - self.layer_flops_per_expert_token = torch.tensor( - [ - # Each expert-token performs gate, up, and down projections; each MAC counts as 2 FLOPs. - 2 * 3 * weight.hidden_size * weight.moe_intermediate_size - for weight in self.weights - ], - dtype=torch.float64, - ) - self.layer_topks = torch.tensor( - [weight.num_experts_per_tok for weight in self.weights], - dtype=torch.float64, - ) - self._round_buffer_storage = array("q", [0]) * (EP_BALANCE_ROUND_BUFFER_CAPACITY * len(self.weights) * 2) - self._round_buffer = torch.frombuffer(self._round_buffer_storage, dtype=torch.int64).view( - EP_BALANCE_ROUND_BUFFER_CAPACITY, len(self.weights), 2 - ) - self._round_ready = threading.Event() - self._written_round_count = 0 # Prefill rounds fully written to the ring buffer. - self._processed_round_count = 0 # Prefill rounds consumed by the monitor thread. - self._overflowed = False - self._previous_pressure_signature: Optional[torch.Tensor] = None - self._common_round_end = torch.zeros((), dtype=torch.int64) - - self.gloo_group = dist_group_manager.ep_balance_monitor_group - if self.gloo_group is None: - raise RuntimeError("EP balance monitor requires a pre-created dedicated Gloo process group") - self.metric_client = MetricClient(get_shm_port_args().metric_port) if self.global_rank == 0 else None - threading.Thread(target=self._monitor_loop, daemon=True, name="ep-balance-monitor").start() - - def record_prefill_round(self): - """Publish one complete all-layer prefill sample to the SPSC ring.""" - if not self.enabled: - return - - written_round_count = self._written_round_count - if written_round_count - self._processed_round_count >= EP_BALANCE_ROUND_BUFFER_CAPACITY: - if not self._overflowed: - self._overflowed = True - self._round_ready.set() - return - - storage_index = (written_round_count % EP_BALANCE_ROUND_BUFFER_CAPACITY) * len(self.counters) * 2 - for counter in self.counters: - self._round_buffer_storage[storage_index] = counter.route_load - self._round_buffer_storage[storage_index + 1] = counter.compute_load - counter.route_load = 0 - counter.compute_load = 0 - storage_index += 2 - - # Publish only after the entire slot is written. The SPSC producer and - # monitor thread run under the CPython GIL, so this count is the release - # point for the corresponding ring slot. - self._written_round_count = written_round_count + 1 - if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - self._round_ready.set() - - def _get_common_round_end(self) -> int: - """Return the exclusive round boundary completed by every rank.""" - self._common_round_end.fill_(self._written_round_count) - dist.all_reduce(self._common_round_end, op=dist.ReduceOp.MIN, group=self.gloo_group) - return int(self._common_round_end.item()) - - def _raise_buffer_overflow(self, phase: str, common_round_end: Optional[int] = None): - message = ( - "EP balance prefill-round buffer overflowed " - f"phase={phase} written={self._written_round_count} " - f"processed={self._processed_round_count} capacity={EP_BALANCE_ROUND_BUFFER_CAPACITY}" - ) - if common_round_end is not None: - message += f" common_round_end={common_round_end}" - raise RuntimeError(message) - - def _copy_local_rounds(self, start: int, end: int) -> torch.Tensor: - """Copy local prefill-round loads in the half-open range [start, end).""" - num_rounds = end - start - if num_rounds > EP_BALANCE_ROUND_BUFFER_CAPACITY: - raise ValueError("requested EP balance round range exceeds ring capacity") - start_index = start % EP_BALANCE_ROUND_BUFFER_CAPACITY - if start_index + num_rounds <= EP_BALANCE_ROUND_BUFFER_CAPACITY: - return self._round_buffer[start_index : start_index + num_rounds].clone() - end_index = (start_index + num_rounds) % EP_BALANCE_ROUND_BUFFER_CAPACITY - return torch.cat((self._round_buffer[start_index:], self._round_buffer[:end_index]), dim=0) - - def _gather_round_stats(self, local_round_stats: torch.Tensor) -> Optional[torch.Tensor]: - """Gather rank-local stats as [round, layer, rank, route/compute].""" - gathered = ( - [torch.empty_like(local_round_stats) for _ in range(self.world_size)] if self.global_rank == 0 else None - ) - dist.gather(local_round_stats, gather_list=gathered, dst=0, group=self.gloo_group) - if self.global_rank != 0: - return None - # [rank, round, layer, route/compute] - # -> [round, layer, rank, route/compute] - return torch.stack(gathered).permute(1, 2, 0, 3) - - def _log_stats(self, round_stats: torch.Tensor): - """Compute and log balance statistics for one complete global window.""" - compute = calculate_prefill_balance_stats( - round_stats, - self.layer_routed_experts, - self.layer_flops_per_expert_token, - self.layer_topks, - self.source_token_replication, - ) - if compute is None: - return - - drift, self._previous_pressure_signature = calculate_prefill_placement_pressure_drift( - round_stats, - previous_pressure_signature=self._previous_pressure_signature, - ) - drift_state = classify_prefill_placement_pressure_drift(drift) - - logger.info( - "ep_balance " - f"phase=prefill prefill_rounds={compute['prefill_rounds']} " - "prefill_ep_critical_overhead_gflops_per_routed_token=" - f"{compute['critical_overhead_gflops_per_routed_token']:.4f} " - "prefill_ep_compute_critical_overhead_ratio=" - f"{compute['prefill_ep_compute_critical_overhead_ratio']:.4f} " - f"prefill_ep_placement_pressure_drift={drift:.4f} " - f"prefill_ep_placement_pressure_state={drift_state}" - ) - self.metric_client.gauge_set( - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", - compute["critical_overhead_gflops_per_routed_token"], - ) - self.metric_client.gauge_set( - "lightllm_prefill_ep_compute_critical_overhead_ratio", - compute["prefill_ep_compute_critical_overhead_ratio"], - ) - self.metric_client.gauge_set("lightllm_prefill_ep_placement_pressure_drift", drift) - - def _monitor_loop(self): - """Consume commonly completed rounds in background report-sized windows.""" - try: - while True: - self._round_ready.wait() - self._round_ready.clear() - if self._overflowed: - self._raise_buffer_overflow("before_sync") - common_round_end = self._get_common_round_end() - if common_round_end - self._processed_round_count > EP_BALANCE_ROUND_BUFFER_CAPACITY: - self._raise_buffer_overflow("common_round_lag", common_round_end=common_round_end) - - while common_round_end - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - round_start = self._processed_round_count - round_end = round_start + EP_BALANCE_PREFILL_ROUNDS_PER_REPORT - local_round_stats = self._copy_local_rounds(round_start, round_end) - if self._overflowed: - self._raise_buffer_overflow("after_copy") - round_stats = self._gather_round_stats(local_round_stats) - self._processed_round_count = round_end - if self.global_rank == 0: - self._log_stats(round_stats) - - if self._written_round_count - self._processed_round_count >= EP_BALANCE_PREFILL_ROUNDS_PER_REPORT: - self._round_ready.set() - except Exception as exc: - logger.exception(f"EP balance monitor stopped unexpectedly: {exc}") - self._disable() - return - - def _disable(self): - """Detach counters from MoE weights and disable monitoring.""" - for weight in self.weights: - weight.fuse_moe_impl.ep_balance_counters = None - self.enabled = False diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py b/lightllm/server/router/model_infer/mode_backend/eplb_manager.py deleted file mode 100644 index 456565e81d..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/eplb_manager.py +++ /dev/null @@ -1,546 +0,0 @@ -import threading -import time -from typing import Dict, Optional - -import torch -import torch.distributed as dist - -from lightllm.common.basemodel.basemodel import TpPartBaseModel -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_logical_to_physical_maps_for_layers, - plan_redundant_experts, - select_improving_placements, -) -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( - NixlEPLBTransfer, - align_target_placement, - build_transfer_plan, -) -from lightllm.utils.dist_utils import get_global_rank, get_global_world_size, get_node_world_size -from lightllm.utils.envs_utils import ( - get_eplb_placement_stickiness, - get_eplb_rebalance_gain_threshold, - get_prefill_eplb_step_interval, -) -from lightllm.utils.log_utils import init_logger - -logger = init_logger(__name__) -EPLB_MIN_AVG_TOKENS_PER_EXPERT = 100 -EPLB_EXPERT_ALIGNMENT = 128 -EPLB_CONTROL_ERROR = -1 -EPLB_STEADY_SAMPLE_STEPS = 4 - - -class EPLBManager: - """Online EPLB with asynchronous GPU expert migration.""" - - def __init__(self, model: TpPartBaseModel): - self.weights = _find_fused_moe_weights(model) - assert self.weights, "EPLB requires at least one EP MoE layer" - self.global_rank = get_global_rank() - self.world_size = get_global_world_size() - self.node_world_size = get_node_world_size() - self._eplb_states = [weight.expert_parallel_state.eplb for weight in self.weights] - self.step_interval = get_prefill_eplb_step_interval() - self.rebalance_gain_threshold = get_eplb_rebalance_gain_threshold() - self.placement_stickiness = get_eplb_placement_stickiness() - self.sampling_interval = self.step_interval - self.prefill_steps = 0 - routed = {weight.expert_parallel_state.num_logical_experts for weight in self.weights} - redundant = {state.num_redundant_experts_per_rank for state in self._eplb_states} - assert len(routed) == len(redundant) == 1 - self.num_logical_experts = routed.pop() - self.num_redundant_experts_per_rank = redundant.pop() - self.current_placement = torch.stack( - [state.initial_redundant_expert_ids_by_rank for state in self._eplb_states] - ) - self.in_flight = False - self.target_placement = None - self.target_metadata = None - self.in_flight_started_at = None - self.evaluation_in_flight = False - self._evaluation_lock = threading.Lock() - self._evaluation_result = None - self._evaluation_error = None - self._evaluation_thread = None - # A fresh manager starts with one continuous base window. After a - # sufficient evaluation, steady state returns to the cheap sparse - # probe. An insufficient sparse probe schedules one fresh continuous - # base window before the next fixed sampling boundary. - self._continuous_collection_start_step: Optional[int] = None - self._continuous_collection_end_step: Optional[int] = self.step_interval - self._steady_collection_end_step: Optional[int] = None - self._reset_recorded_samples() - self._set_recording(True) - # Keep background evaluation collectives separate from the main-thread - # control/poll collectives: their ordering is intentionally independent. - self.evaluation_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self.control_group = dist.new_group(list(range(self.world_size)), backend="gloo") - self.transfer_group = dist.new_group(list(range(self.world_size)), backend="gloo") - # This control-group scalar is only touched from the main inference - # thread, never by the background evaluation thread. - self._control_ready_count = torch.empty(1, dtype=torch.int32) - self.transfer = NixlEPLBTransfer(self.weights, self.transfer_group, self.global_rank, self.world_size) - if self.global_rank == 0: - logger.info( - "eplb enabled " - f"layers={len(self.weights)} num_logical_experts={self.num_logical_experts} " - f"num_redundant_experts_per_rank={self.num_redundant_experts_per_rank} " - f"step_interval={self.step_interval} " - f"rebalance_gain_threshold={self.rebalance_gain_threshold:.4f} " - f"placement_stickiness={self.placement_stickiness:.4f}" - ) - - def poll(self): - """Poll only from a globally ordered pre-forward boundary.""" - if self.in_flight: - self._poll_in_flight() - return - if self.evaluation_in_flight and self._evaluation_ready_on_all_ranks(): - self._poll_evaluation() - - def step(self): - if self.in_flight or self.evaluation_in_flight: - return - self.prefill_steps += 1 - continuous_start = self._continuous_collection_start_step - continuous_end = self._continuous_collection_end_step - if continuous_end is not None: - if continuous_start is not None and self.prefill_steps == continuous_start: - self._set_recording(True) - if self.prefill_steps >= continuous_end: - self._start_evaluation() - return - sampling_interval = self.sampling_interval - phase = self.prefill_steps % sampling_interval - steady_collection_end_step = self._steady_collection_end_step - if steady_collection_end_step is not None: - if self.prefill_steps >= steady_collection_end_step: - self._steady_collection_end_step = None - self._start_evaluation() - return - if sampling_interval == 1: - self._start_evaluation() - return - if phase == sampling_interval - self._steady_sample_window_steps(): - self._arm_steady_collection(self.prefill_steps + self._steady_sample_window_steps()) - - def _set_recording(self, enabled: bool): - for state in self._eplb_states: - state.recording = enabled - - def _reset_recorded_samples(self): - counters = [state.route_counter for state in self._eplb_states] - if counters: - torch._foreach_zero_(counters) - for state in self._eplb_states: - state.recorded_sample_count = 0 - - def _control_count(self, value: int) -> torch.Tensor: - """Return the main-thread-only reusable control collective scalar.""" - return self._control_ready_count.fill_(value) - - def _clear_continuous_collection(self): - self._continuous_collection_start_step = None - self._continuous_collection_end_step = None - - def _steady_sample_window_steps(self) -> int: - return min(EPLB_STEADY_SAMPLE_STEPS, self.sampling_interval) - - def _arm_steady_collection(self, collection_end_step: int): - """Start the fixed sparse window without moving its evaluation boundary.""" - self._reset_recorded_samples() - self._steady_collection_end_step = collection_end_step - self._set_recording(True) - - def _begin_continuous_collection(self): - minimum_end = self.prefill_steps + self.step_interval - collection_end = -(-minimum_end // self.sampling_interval) * self.sampling_interval - self._reset_recorded_samples() - self._steady_collection_end_step = None - self._continuous_collection_start_step = collection_end - self.step_interval - self._continuous_collection_end_step = collection_end - self._set_recording(self._continuous_collection_start_step == self.prefill_steps) - - def _prepare_next_sampling_window(self): - """Clear the current window and arm the next sparse sampling window.""" - self._clear_continuous_collection() - self._steady_collection_end_step = None - if self.sampling_interval == 1: - self._reset_recorded_samples() - self._set_recording(True) - elif self.sampling_interval <= EPLB_STEADY_SAMPLE_STEPS: - # There is no later pre-boundary manager step at which to arm a - # full clamped window, so arm immediately but keep the same next - # fixed boundary. - self._arm_steady_collection(self.prefill_steps + self.sampling_interval) - else: - self._reset_recorded_samples() - self._set_recording(False) - - @staticmethod - def _recent_ring_samples(counter: torch.Tensor, recorded_sample_count: int) -> torch.Tensor: - """Return the newest ring rows in chronological order.""" - capacity = counter.shape[0] - available = min(recorded_sample_count, capacity) - if available == 0: - return counter[:0] - start = (recorded_sample_count - available) % capacity - indices = (torch.arange(available, dtype=torch.int64, device=counter.device) + start) % capacity - return counter.index_select(0, indices) - - def _collect_local_samples(self) -> torch.Tensor: - counters = [state.route_counter for state in self._eplb_states] - capacities = [counter.shape[0] for counter in counters] - if len(set(capacities)) != 1 or any(counter.ndim != 2 for counter in counters): - raise RuntimeError("EPLB sample capacities differ between layers") - counts = [state.recorded_sample_count for state in self._eplb_states] - if len(set(counts)) != 1: - raise RuntimeError("EPLB recorded sample counts differ between layers") - sample_count = counts[0] - # Validate the metadata before copying the newest rows to the CPU. - metadata = torch.tensor([sample_count, -sample_count, capacities[0], -capacities[0]], dtype=torch.int64) - dist.all_reduce(metadata, op=dist.ReduceOp.MIN, group=self.evaluation_group) - if metadata[0] != -metadata[1] or metadata[2] != -metadata[3]: - raise RuntimeError("EPLB recorded sample count or capacity differs between ranks") - # Stack the fixed-size ring buffers in one GPU launch. Slicing each - # layer before stacking turns a single launch into one index_select per - # MoE layer and is measurably slower in the normal sparse path. - counter_samples = torch.stack(counters, dim=1) - return self._recent_ring_samples(counter_samples, sample_count).cpu() - - def _commit_layer_metadata(self, layer_index: int): - eplb_state = self._eplb_states[layer_index] - logical_to_physical, replica_count = self.target_metadata[layer_index] - eplb_state.logical_to_physical_map.copy_(logical_to_physical, non_blocking=True) - eplb_state.logical_replica_count.copy_(replica_count, non_blocking=True) - - def _finish_rebalance(self): - self.current_placement = self.target_placement - self.target_placement = None - self.target_metadata = None - self.in_flight = False - self._prepare_next_sampling_window() - if self.global_rank == 0: - logger.info(f"eplb completed wall_time={time.time() - self.in_flight_started_at:.2f}s") - - def _poll_in_flight(self): - local_error = None - try: - pending = self.transfer.pending_layers() - except BaseException as exc: - pending = [] - local_error = exc - ready_count = self._control_count(EPLB_CONTROL_ERROR if local_error is not None else len(pending)) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: - if local_error is not None: - raise RuntimeError("EPLB transfer worker failed on this rank") from local_error - raise RuntimeError("EPLB transfer worker failed on another rank") - if ready_count == 0: - return - if ready_count > len(pending) or ready_count > len(self.in_flight_layers): - raise RuntimeError("EPLB global ready count exceeds the local ordered prefix") - from lightllm.server.router.model_infer.infer_batch import g_infer_context - - # Previous forward is queued on the shared overlap stream; order the - # live-weight commit after it. The subsequent wait orders the next forward. - torch.cuda.current_stream().wait_stream(g_infer_context.get_overlap_stream()) - for layer_index, buffer_index in pending[:ready_count]: - if layer_index != self.in_flight_layers[0]: - raise RuntimeError( - f"EPLB pending layer {layer_index} does not match expected {self.in_flight_layers[0]}" - ) - self.transfer.commit(layer_index, buffer_index, lambda: self._commit_layer_metadata(layer_index)) - self.in_flight_layers.pop(0) - if not self.in_flight_layers: - self.transfer.finish() - self._finish_rebalance() - - def _plan_and_broadcast(self, global_load: torch.Tensor): - """Plan on rank zero and share the serializable result on the evaluation group.""" - result = None - local_error = None - if self.global_rank == 0: - try: - minimum = self.num_logical_experts * EPLB_MIN_AVG_TOKENS_PER_EXPERT - layer_samples = global_load.sum(dim=(0, 2, 3)) - if torch.any(layer_samples < minimum): - result = { - "kind": "insufficient", - "minimum_layer_samples": int(layer_samples.min().item()), - "minimum": minimum, - } - else: - candidate = plan_redundant_experts( - global_load, - self.world_size, - self.num_redundant_experts_per_rank, - expert_alignment=EPLB_EXPERT_ALIGNMENT, - node_world_size=self.node_world_size, - current_placement=self.current_placement, - stickiness=self.placement_stickiness, - ) - placement, improved, metrics, before_load, after_load = select_improving_placements( - global_load, - self.current_placement, - candidate, - expert_alignment=EPLB_EXPERT_ALIGNMENT, - node_world_size=self.node_world_size, - rebalance_gain_threshold=self.rebalance_gain_threshold, - ) - if bool(torch.any(improved)): - # A planner placement identifies experts by rank, not - # by redundant slot. Canonicalize every selected row - # before broadcasting so transfer, metadata, and the - # next current_placement all describe the same live - # physical expert rows. - placement = placement.clone() - for layer_index in torch.nonzero(improved, as_tuple=False).flatten().tolist(): - placement[layer_index] = align_target_placement( - self.current_placement[layer_index], placement[layer_index] - ) - result = { - "kind": "planned" if bool(torch.any(improved)) else "no_improvement", - "placement": placement, - "improved": improved, - "before": _imbalance_summary(before_load), - "after": _imbalance_summary(after_load), - **metrics, - } - except BaseException as exc: - local_error = exc - result = {"kind": "error", "message": f"{type(exc).__name__}: {exc}"} - if self.world_size > 1: - result_list = [result] - dist.broadcast_object_list(result_list, src=0, group=self.evaluation_group) - result = result_list[0] - if result["kind"] == "error": - if local_error is not None: - raise RuntimeError("EPLB planner failed on rank zero") from local_error - raise RuntimeError(f"EPLB planner failed on rank zero: {result['message']}") - return result - - def _evaluate_after_event(self, event: torch.cuda.Event): - """Run the CPU/Gloo planning phase after the frozen CUDA counters are ready.""" - try: - torch.cuda.set_device(self._eplb_states[0].route_counter.device) - event.synchronize() - local_load = self._collect_local_samples() - recorded_sample_count = int(local_load.shape[0]) - sample_window_steps = ( - self.step_interval - if self._continuous_collection_end_step is not None - else self._steady_sample_window_steps() - ) - num_nodes = self.world_size // self.node_world_size - # Preserve source nodes until physical-replica loads are combined; - # DeepEP applies expert alignment after traffic from all sources - # reaches each destination expert. - global_load = torch.zeros((*local_load.shape[:2], num_nodes, local_load.shape[2]), dtype=local_load.dtype) - global_load[:, :, self.global_rank // self.node_world_size] = local_load - dist.all_reduce(global_load, op=dist.ReduceOp.SUM, group=self.evaluation_group) - result = self._plan_and_broadcast(global_load) - result["recorded_sample_count"] = recorded_sample_count - result["sample_window_steps"] = sample_window_steps - if result["kind"] == "planned": - metadata = [None] * len(self.weights) - layer_plans = [] - improved_layer_indices = torch.nonzero(result["improved"], as_tuple=False).flatten() - if improved_layer_indices.numel(): - maps_for_improved_layers, counts_for_improved_layers = build_logical_to_physical_maps_for_layers( - result["placement"][improved_layer_indices], - self.num_logical_experts, - source_rank=self.global_rank, - node_world_size=self.node_world_size, - ) - for improved_layer_offset, layer_index in enumerate(improved_layer_indices.tolist()): - placement = result["placement"][layer_index] - metadata[layer_index] = ( - maps_for_improved_layers[improved_layer_offset], - counts_for_improved_layers[improved_layer_offset], - ) - layer_plans.append( - ( - layer_index, - build_transfer_plan( - self.current_placement[layer_index], - placement, - self.num_logical_experts, - self.world_size, - self.node_world_size, - ), - ) - ) - result["metadata"] = metadata - result["layer_plans"] = layer_plans - result["prepared_batches"] = self.transfer.prepare_transfer(layer_plans) - with self._evaluation_lock: - self._evaluation_result = result - except BaseException as exc: - with self._evaluation_lock: - self._evaluation_error = exc - - def _start_evaluation(self): - with self._evaluation_lock: - self._evaluation_result = None - self._evaluation_error = None - self._set_recording(False) - event = torch.cuda.Event() - event.record(torch.cuda.current_stream()) - self.evaluation_in_flight = True - self._evaluation_thread = threading.Thread(target=self._evaluate_after_event, args=(event,), daemon=True) - self._evaluation_thread.start() - - def _poll_evaluation(self): - if not self.evaluation_in_flight: - return False - with self._evaluation_lock: - error = self._evaluation_error - result = self._evaluation_result - if error is not None or result is not None: - self._evaluation_result = None - self._evaluation_error = None - if error is not None: - self._evaluation_thread.join() - self.evaluation_in_flight = False - self._evaluation_thread = None - raise error - if result is None: - return True - self._evaluation_thread.join() - self.evaluation_in_flight = False - self._evaluation_thread = None - if result["kind"] == "insufficient": - from_continuous_window = self._continuous_collection_end_step is not None - if from_continuous_window: - self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) - self._prepare_next_sampling_window() - else: - self._begin_continuous_collection() - if self.global_rank == 0: - if from_continuous_window: - logger.info( - "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " - "next_sampling_interval=%s recorded_sample_count=%s sample_window_steps=%s", - self.prefill_steps, - result["minimum_layer_samples"], - result["minimum"], - self.sampling_interval, - result.get("recorded_sample_count"), - result.get("sample_window_steps"), - ) - else: - logger.info( - "eplb insufficient samples: prefill_steps=%s minimum_layer_samples=%s required=%s " - "scheduled_fresh_window_start=%s scheduled_fresh_window_end=%s " - "recorded_sample_count=%s sample_window_steps=%s", - self.prefill_steps, - result["minimum_layer_samples"], - result["minimum"], - self._continuous_collection_start_step, - self._continuous_collection_end_step, - result.get("recorded_sample_count"), - result.get("sample_window_steps"), - ) - return False - if result["kind"] == "no_improvement": - self.sampling_interval = min(self.sampling_interval * 4, self.step_interval * 16) - if self.global_rank == 0: - logger.info( - "eplb skip rearrangement: no model improvement model_imbalance_ratio=%.4f " - "candidate_model_imbalance_ratio=%.4f candidate_rebalance_gain=%.4f " - "candidate_changed_layer_count=%s actual_changed_layer_count=0 next_sampling_interval=%s " - "recorded_sample_count=%s sample_window_steps=%s", - result["model_imbalance_ratio"], - result["candidate_model_imbalance_ratio"], - result["candidate_rebalance_gain"], - result["candidate_changed_layer_count"], - self.sampling_interval, - result.get("recorded_sample_count"), - result.get("sample_window_steps"), - ) - self._prepare_next_sampling_window() - return False - self._start_rebalance(result) - return True - - def _evaluation_ready_on_all_ranks(self) -> bool: - with self._evaluation_lock: - local_error = self._evaluation_error - local_result = self._evaluation_result - local_status = EPLB_CONTROL_ERROR if local_error is not None else int(local_result is not None) - ready_count = self._control_count(local_status) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=self.control_group) - ready_count = int(ready_count.item()) - if ready_count < 0: - if local_error is not None: - raise RuntimeError("EPLB evaluation failed on this rank") from local_error - raise RuntimeError("EPLB evaluation failed on another rank") - return bool(ready_count) - - def _start_rebalance(self, result): - placement = result["placement"] - layer_plans = result["layer_plans"] - self.sampling_interval = self.step_interval - self._clear_continuous_collection() - self._reset_recorded_samples() - self.target_placement = placement - self.target_metadata = result["metadata"] - self.in_flight_layers = [layer_index for layer_index, _ in layer_plans] - self.in_flight = True - self.in_flight_started_at = time.time() - self.transfer.start(layer_plans, result["prepared_batches"]) - if self.global_rank == 0: - actual_changed_slot_count = sum(len(plan) for _, plan in layer_plans) - cross_node_transfer_count = sum( - step.src_rank // self.node_world_size != step.dst_rank // self.node_world_size - for _, plan in layer_plans - for step in plan - ) - logger.info( - "eplb started prefill_steps=%s max_before=%.4f max_after=%.4f p95_before=%.4f p95_after=%.4f " - "model_imbalance_ratio=%.4f candidate_model_imbalance_ratio=%.4f " - "candidate_rebalance_gain=%.4f candidate_changed_layer_count=%s " - "actual_changed_layer_count=%s actual_changed_slot_count=%s cross_node_transfer_count=%s " - "recorded_sample_count=%s sample_window_steps=%s", - self.prefill_steps, - result["before"]["max"], - result["after"]["max"], - result["before"]["p95"], - result["after"]["p95"], - result["model_imbalance_ratio"], - result["candidate_model_imbalance_ratio"], - result["candidate_rebalance_gain"], - result["candidate_changed_layer_count"], - len(layer_plans), - actual_changed_slot_count, - cross_node_transfer_count, - result.get("recorded_sample_count"), - result.get("sample_window_steps"), - ) - - -def _imbalance_summary(rank_load: torch.Tensor) -> Dict[str, float]: - if rank_load.ndim != 3: - raise ValueError("rank_load must be [samples, layers, ranks]") - critical = rank_load.max(dim=2).values.sum(dim=0) - mean = rank_load.mean(dim=2).sum(dim=0) - layer_imbalance = critical / mean.clamp_min(1.0) - sorted_imbalance = torch.sort(layer_imbalance).values - p95_index = max(0, (95 * layer_imbalance.numel() + 99) // 100 - 1) - return { - "max": float(layer_imbalance.max().item()), - "p95": float(sorted_imbalance[p95_index].item()), - } - - -def _find_fused_moe_weights(model): - weights_by_id = {} - for layer in model.trans_layers_weight: - for value in getattr(layer, "__dict__", {}).values(): - if isinstance(value, FusedMoeWeight) and value.enable_ep_moe: - weights_by_id[id(value)] = value - return sorted(weights_by_id.values(), key=lambda weight: weight.layer_num_) diff --git a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py b/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py deleted file mode 100644 index e5737fcb7a..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/eplb_transfer.py +++ /dev/null @@ -1,759 +0,0 @@ -"""Asynchronous expert-row migration for EPLB.""" -import ctypes -import os -import re -import socket -import threading -from collections import defaultdict, deque -from dataclasses import dataclass -from typing import Dict, Iterable, List, Optional, Sequence, Tuple - -import torch -import torch.distributed as dist - -from lightllm.common.eplb_utils import EPLB_MAX_STAGING_DEPTH, extract_eplb_expert_tensors - - -@dataclass(frozen=True) -class TransferStep: - dst_rank: int - dst_slot: int - src_rank: int - src_local_row: int - - -def align_target_placement(current: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - """Canonicalize a target row layout without moving retained experts. - - EPLB placement is rank-based: redundant slots on one rank are - interchangeable. Retained experts therefore keep their live physical - slot, while new experts fill freed slots in the planner's target-row - order. The returned placement is the single canonical layout that must - be used both for transfers and for published routing metadata. - """ - assert current.ndim == target.ndim == 2 - assert tuple(current.shape) == tuple(target.shape) - - current_rows = current.tolist() - target_rows = target.tolist() - aligned_target_rows = [] - for current_row, target_row in zip(current_rows, target_rows): - current_slots = {expert: slot for slot, expert in enumerate(current_row)} - target_experts = set(target_row) - aligned_row = list(current_row) - freed_slots = [slot for slot, expert in enumerate(current_row) if expert not in target_experts] - new_experts = [expert for expert in target_row if expert not in current_slots] - assert len(freed_slots) == len(new_experts) - for slot, expert in zip(freed_slots, new_experts): - aligned_row[slot] = expert - aligned_target_rows.append(aligned_row) - return target.new_tensor(aligned_target_rows) - - -def build_transfer_plan( - current: torch.Tensor, - target: torch.Tensor, - num_logical_experts: int, - world_size: int, - node_world_size: int, -) -> List[TransferStep]: - assert tuple(current.shape) == tuple(target.shape) == (world_size, current.shape[1]) - num_experts_per_rank = num_logical_experts // world_size - current_rows = current.tolist() - aligned_target_rows = align_target_placement(current, target).tolist() - # A logical expert has one primary row and at most one redundant row per - # rank, so this source list is already unique. Build it once instead of - # allocating/sorting a set for every destination slot. - candidates_by_expert = [ - [ - ( - expert // num_experts_per_rank, - expert % num_experts_per_rank, - ) - ] - for expert in range(num_logical_experts) - ] - for rank, row in enumerate(current_rows): - for slot, expert in enumerate(row): - candidates_by_expert[expert].append((rank, num_experts_per_rank + slot)) - source_load = [0] * world_size - plan = [] - for dst_rank in range(world_size): - for dst_slot, expert in enumerate(aligned_target_rows[dst_rank]): - if expert == current_rows[dst_rank][dst_slot]: - continue - src_rank, src_row = min( - candidates_by_expert[expert], - key=lambda item: ( - item[0] // node_world_size != dst_rank // node_world_size, - source_load[item[0]], - item[0], - item[1], - ), - ) - source_load[src_rank] += 1 - plan.append(TransferStep(dst_rank, dst_slot, src_rank, src_row)) - return plan - - -class _EPLBTransferBase: - """Shared live/staging buffers and publish/commit lifecycle.""" - - def __init__(self, weights, transfer_group, global_rank, world_size): - self._eplb_states = [weight.expert_parallel_state.eplb for weight in weights] - self.transfer_group = transfer_group - self.global_rank = global_rank - self.world_size = world_size - self.num_experts_per_rank = weights[0].expert_parallel_state.num_primary_experts_per_rank - self.device = weights[0].w13.weight.device - self.live = [extract_eplb_expert_tensors(weight) for weight in weights] - self._validate_live_layout() - num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank - self.staging = [ - [ - ( - name, - torch.empty( - (num_redundant_slots_per_rank,) + tuple(tensor.shape[1:]), - dtype=tensor.dtype, - device=tensor.device, - ), - ) - for name, tensor in self.live[0] - ] - for _ in range(self.staging_depth) - ] - self._release = [threading.Event() for _ in range(self.staging_depth)] - for release in self._release: - release.set() - self._error = None - self._consumed_events = [torch.cuda.Event() for _ in range(self.staging_depth)] - self._consumed_recorded = [False] * self.staging_depth - self._changed_dst_slots = [()] * self.staging_depth - self._pending = deque() - self._pending_lock = threading.Lock() - self._thread = None - self._needs_staging_reuse_barrier = False - - def _validate_live_layout(self) -> None: - reference = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in self.live[0]] - num_redundant_slots_per_rank = self._eplb_states[0].num_redundant_experts_per_rank - for layer_index, (state, tensors) in enumerate(zip(self._eplb_states, self.live)): - layout = [(name, tuple(tensor.shape[1:]), tensor.dtype, tensor.device) for name, tensor in tensors] - assert layout == reference, f"EPLB layer {layer_index} has incompatible expert tensor layout" - assert ( - state.num_redundant_experts_per_rank == num_redundant_slots_per_rank - ), "EPLB redundant slot count must match" - - def _make_batches(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [ - [ - (layer_index, plan, buffer_index, self.staging[buffer_index]) - for buffer_index, (layer_index, plan) in enumerate( - layer_plans[batch_start : batch_start + self.staging_depth] - ) - ] - for batch_start in range(0, len(layer_plans), self.staging_depth) - ] - - def start(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]], prepared_batches) -> None: - if self._thread is not None and self._thread.is_alive(): - raise RuntimeError("EPLB transfer is already in flight") - expected_batch_count = (len(layer_plans) + self.staging_depth - 1) // self.staging_depth - if len(prepared_batches) != expected_batch_count: - raise ValueError("EPLB prepared batch count does not match layer-plan batches") - # 预构造批次与描述符一起传入,避免推理线程重建。 - self._start_transfer_generation() - self._error = None - with self._pending_lock: - self._pending.clear() - - def worker() -> None: - try: - torch.cuda.set_device(self.device) - if not layer_plans: - self._finish_transfer_generation() - for batch_index, (batch, prepared_batch) in enumerate(prepared_batches): - batch_start = batch_index * self.staging_depth - for layer_index, plan, buffer_index, _ in batch: - release = self._release[buffer_index] - # A buffer cannot be reused until its prior committed rows are no longer read by CUDA. - release.wait() - release.clear() - if self._consumed_recorded[buffer_index]: - self._consumed_events[buffer_index].synchronize() - self._changed_dst_slots[buffer_index] = tuple( - step.dst_slot for step in plan if step.dst_rank == self.global_rank - ) - if batch_start > 0 and self._needs_staging_reuse_barrier: - # All destinations must finish consuming the prior IPC staging generation - # before a source can reuse the peer buffer for this batch. - dist.barrier(group=self.transfer_group) - self._copy_batch(batch, prepared_batch) - if batch_start + self.staging_depth >= len(layer_plans): - self._finish_transfer_generation() - with self._pending_lock: - self._pending.extend((layer_index, buffer_index) for layer_index, _, buffer_index, _ in batch) - except BaseException as exc: - self._error = exc - - self._thread = threading.Thread(target=worker, name=f"eplb-{self.backend}", daemon=True) - self._thread.start() - - def pending_layers(self): - if self._error is not None: - raise RuntimeError("EPLB migration worker failed") from self._error - with self._pending_lock: - return list(self._pending) - - def commit(self, layer_index: int, buffer_index: int, post_copy=None) -> None: - with self._pending_lock: - if not self._pending or self._pending[0] != (layer_index, buffer_index): - raise RuntimeError("EPLB commit does not match the pending FIFO") - self._pending.popleft() - changed_dst_slots = self._changed_dst_slots[buffer_index] - for (_, live), (_, staging) in zip(self.live[layer_index], self.staging[buffer_index]): - _commit_staging_rows( - live, - staging, - self.num_experts_per_rank, - changed_dst_slots, - ) - if post_copy is not None: - post_copy() - self._consumed_events[buffer_index].record(torch.cuda.current_stream()) - self._consumed_recorded[buffer_index] = True - self._release[buffer_index].set() - - def finish(self) -> None: - """Wait for the released migration worker to exit before another rebalance.""" - thread = self._thread - if thread is None: - return - thread.join() - self._thread = None - if self._error is not None: - raise RuntimeError("EPLB migration worker failed") from self._error - - -class NixlEPLBTransfer(_EPLBTransferBase): - """GPU-direct UCX/NIXL EPLB transfer. Initialization errors are fatal.""" - - backend = "nixl" - _DEFAULT_UCX_TLS = "self,sm,cuda_ipc,cuda_copy,rc_x" - - @dataclass - class _PreparedBatch: - remote_entries: Dict[int, list] - push_batch: "_PreparedCudaMemcpyBatch | None" - - def __init__(self, weights, transfer_group, global_rank, world_size): - # Reuse at most eight layer buffers to bound EPLB staging memory. - self.staging_depth = min(EPLB_MAX_STAGING_DEPTH, len(weights)) - super().__init__(weights, transfer_group, global_rank, world_size) - self._nixl_agent = None - self._registered_descs = None - self._remote_agents: Dict[int, str] = {} - self._remote_layouts = {} - self._xfer_cache = {} - self._used_xfer_cache_keys = set() - self._ipc_staging = {} - self._same_node_ranks = set() - self._cross_node_ranks = set() - self._push_stream = torch.cuda.Stream(device=self.device) - self._batch_memcpy = _CudaBatchMemcpy() - try: - self._init_ipc_metadata() - self._init_push_layouts() - if self._cross_node_ranks: - os.environ.setdefault("UCX_TLS", self._DEFAULT_UCX_TLS) - try: - import nixl - except Exception as exc: - raise RuntimeError("NIXL EPLB backend requires the nixl package for cross-node transfer") from exc - agent_name = f"lightllm-eplb-{socket.gethostname()}-{os.getpid()}-rank-{global_rank}" - config = nixl.nixl_agent_config(enable_prog_thread=True, enable_listen_thread=False, backends=["UCX"]) - self._nixl_agent = nixl.nixl_agent(agent_name, config) - reg_tensors = [tensor for layer in self.live for _, tensor in layer] + [ - tensor for staging in self.staging for _, tensor in staging - ] - self._registered_descs = self._nixl_agent.get_reg_descs(reg_tensors) - self._nixl_agent.register_memory(self._registered_descs, backends=["UCX"]) - self._init_remote_metadata() - except Exception as exc: - self.shutdown() - if isinstance(exc, RuntimeError): - raise - raise RuntimeError("NIXL EPLB initialization failed") from exc - - def _local_layout(self): - return [ - [(name, tensor.data_ptr(), tensor.get_device(), tensor[0].nbytes) for name, tensor in layer] - for layer in self.live - ] - - def _init_ipc_metadata(self) -> None: - hostnames = [None] * self.world_size - dist.all_gather_object(hostnames, socket.gethostname(), group=self.transfer_group) - local_hostname = hostnames[self.global_rank] - self._needs_staging_reuse_barrier = len(set(hostnames)) < len(hostnames) - self._same_node_ranks = {rank for rank, hostname in enumerate(hostnames) if hostname == local_hostname} - self._cross_node_ranks = set(range(self.world_size)) - self._same_node_ranks - from lightllm.server.router.model_infer.mode_backend.pd.p2p_fix import ( - p2p_fix_rebuild_cuda_tensor, - reduce_tensor, - ) - - exports = {} - for target_rank in self._same_node_ranks - {self.global_rank}: - exports[target_rank] = { - "staging": [ - [(name, tuple(tensor.shape), tensor.dtype, reduce_tensor(tensor)[1]) for name, tensor in staging] - for staging in self.staging - ], - } - all_exports = [None] * self.world_size - dist.all_gather_object(all_exports, exports, group=self.transfer_group) - - torch.cuda.set_device(self.device) - for dst_rank in self._same_node_ranks - {self.global_rank}: - metadata = all_exports[dst_rank].get(self.global_rank) - if metadata is None or len(metadata["staging"]) != self.staging_depth: - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} has incompatible staging metadata") - rebuilt_staging = [] - for remote_staging, local_staging in zip(metadata["staging"], self.staging): - if len(remote_staging) != len(local_staging): - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging tensor count mismatch") - rebuilt = [] - for (name, shape, dtype, args), (local_name, local_tensor) in zip(remote_staging, local_staging): - if name != local_name or shape != tuple(local_tensor.shape) or dtype != local_tensor.dtype: - raise RuntimeError(f"NIXL IPC destination rank {dst_rank} staging layout mismatch for {name}") - tensor = p2p_fix_rebuild_cuda_tensor(*args) - if tuple(tensor.shape) != shape or tensor.dtype != dtype or tensor.device != local_tensor.device: - raise RuntimeError( - f"NIXL IPC destination rank {dst_rank} staging rebuild validation failed for {name}" - ) - rebuilt.append((name, tensor)) - rebuilt_staging.append(rebuilt) - self._ipc_staging[dst_rank] = rebuilt_staging - - def _init_remote_metadata(self) -> None: - metadata = self._nixl_agent.get_agent_metadata() - all_metadata = [None] * self.world_size - all_layouts = [None] * self.world_size - dist.all_gather_object(all_metadata, metadata, group=self.transfer_group) - dist.all_gather_object(all_layouts, self._local_layout(), group=self.transfer_group) - for rank in self._cross_node_ranks: - layout = all_layouts[rank] - if len(layout) != len(self.live): - raise RuntimeError(f"NIXL remote rank {rank} has incompatible layer layout") - self._remote_agents[rank] = self._nixl_agent.add_remote_agent(all_metadata[rank]) - self._remote_layouts[rank] = layout - - def _wait_xfers(self, xfers) -> None: - pending = [] - for item in xfers: - state = self._nixl_agent.transfer(item[2]) - if state == "ERR": - raise RuntimeError("NIXL READ post failed") - if state == "PROC": - pending.append(item) - while pending: - remaining = [] - for item in pending: - state = self._nixl_agent.check_xfer_state(item[2]) - if state == "ERR": - raise RuntimeError("NIXL READ transfer failed") - if state != "DONE": - remaining.append(item) - pending = remaining - - def _release_xfers(self, xfers) -> None: - unreleased = [] - errors = [] - for local_dlist, remote_dlist, xfer in xfers: - remaining = [local_dlist, remote_dlist, xfer] - for remaining_index, handle, release in ( - (2, xfer, self._nixl_agent.release_xfer_handle), - (1, remote_dlist, self._nixl_agent.release_dlist_handle), - (0, local_dlist, self._nixl_agent.release_dlist_handle), - ): - if handle is not None: - try: - release(handle) - except Exception as exc: - errors.append(exc) - else: - remaining[remaining_index] = None - if any(handle is not None for handle in remaining): - unreleased.append(tuple(remaining)) - if errors: - error = RuntimeError("NIXL transfer handle release failed") - error.unreleased_xfers = unreleased - raise error from errors[0] - - @staticmethod - def _contiguous_runs(steps): - ordered = sorted(steps, key=lambda step: (step.src_local_row, step.dst_slot)) - runs = [] - for step in ordered: - if ( - runs - and step.src_local_row == runs[-1][-1].src_local_row + 1 - and step.dst_slot == runs[-1][-1].dst_slot + 1 - ): - runs[-1].append(step) - else: - runs.append([step]) - return runs - - @staticmethod - def _remote_read_cache_key(src_rank: int, entries): - return ( - src_rank, - tuple( - ( - layer_index, - tuple((step.src_local_row, step.dst_slot) for step in run), - tuple(tensor.data_ptr() for _, tensor in staging), - ) - for layer_index, run, staging in entries - ), - ) - - def _init_push_layouts(self) -> None: - self._live_row_layout = [ - [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in layer] for layer in self.live - ] - reference = [(name, row_nbytes) for name, _, row_nbytes in self._live_row_layout[0]] - self._push_staging_row_layout = {} - for dst_rank in self._same_node_ranks: - layouts = [] - for buffer_index in range(self.staging_depth): - staging = ( - self.staging[buffer_index] - if dst_rank == self.global_rank - else self._ipc_staging[dst_rank][buffer_index] - ) - layout = [(name, tensor.data_ptr(), tensor[0].nbytes) for name, tensor in staging] - if [(name, row_nbytes) for name, _, row_nbytes in layout] != reference: - raise RuntimeError("NIXL source-push staging row layout mismatch") - layouts.append(layout) - self._push_staging_row_layout[dst_rank] = layouts - - def _prepare_batch(self, batch): - remote_entries = defaultdict(list) - push_descriptors = [] - for layer_index, plan, buffer_index, staging in batch: - steps_by_source = defaultdict(list) - by_destination = defaultdict(list) - for step in plan: - if step.dst_rank == self.global_rank and step.src_rank not in self._same_node_ranks: - steps_by_source[step.src_rank].append(step) - if step.src_rank == self.global_rank and step.dst_rank in self._same_node_ranks: - by_destination[step.dst_rank].append(step) - for src_rank, steps in steps_by_source.items(): - remote_entries[src_rank].extend((layer_index, run, staging) for run in self._contiguous_runs(steps)) - source_layout = self._live_row_layout[layer_index] - for dst_rank, steps in by_destination.items(): - destination_layout = self._push_staging_row_layout[dst_rank][buffer_index] - for run in self._contiguous_runs(steps): - first = run[0] - run_len = len(run) - for (_, source_ptr, row_nbytes), (_, destination_ptr, _) in zip(source_layout, destination_layout): - push_descriptors.append( - ( - source_ptr + first.src_local_row * row_nbytes, - destination_ptr + first.dst_slot * row_nbytes, - run_len * row_nbytes, - ) - ) - push_batch = self._batch_memcpy.prepare(push_descriptors) if push_descriptors else None - return self._PreparedBatch(dict(remote_entries), push_batch) - - def prepare_transfer(self, layer_plans: Sequence[Tuple[int, Sequence[TransferStep]]]): - return [(batch, self._prepare_batch(batch)) for batch in self._make_batches(layer_plans)] - - def _get_remote_read(self, src_rank: int, entries): - cache_key = self._remote_read_cache_key(src_rank, entries) - cached = self._xfer_cache.get(cache_key) - if cached is not None: - self._used_xfer_cache_keys.add(cache_key) - return cached - local_descs = [] - remote_descs = [] - local_dlist = remote_dlist = xfer = None - try: - for layer_index, run, staging in entries: - remote_layer = self._remote_layouts[src_rank][layer_index] - if len(remote_layer) != len(staging): - raise RuntimeError(f"NIXL remote rank {src_rank} has incompatible layer layout") - first = run[0] - run_len = len(run) - for tensor_index, (_, staging_tensor) in enumerate(staging): - name, remote_ptr, remote_device, remote_nbytes = remote_layer[tensor_index] - if ( - name != self.live[layer_index][tensor_index][0] - or remote_nbytes != staging_tensor[first.dst_slot].nbytes - ): - raise RuntimeError(f"NIXL remote rank {src_rank} descriptor range mismatch") - local_descs.append( - ( - staging_tensor[first.dst_slot].data_ptr(), - run_len * remote_nbytes, - staging_tensor.get_device(), - ) - ) - remote_descs.append( - (remote_ptr + first.src_local_row * remote_nbytes, run_len * remote_nbytes, remote_device) - ) - local_dlist = self._nixl_agent.prep_xfer_dlist( - "NIXL_INIT_AGENT", self._nixl_agent.get_xfer_descs(local_descs, "VRAM"), backends=["UCX"] - ) - remote_dlist = self._nixl_agent.prep_xfer_dlist( - self._remote_agents[src_rank], self._nixl_agent.get_xfer_descs(remote_descs, "VRAM"), backends=["UCX"] - ) - xfer = self._nixl_agent.make_prepped_xfer( - "READ", - local_dlist, - list(range(len(local_descs))), - remote_dlist, - list(range(len(remote_descs))), - backends=["UCX"], - ) - selected_backend = self._nixl_agent.query_xfer_backend(xfer) - if selected_backend != "UCX": - raise RuntimeError("NIXL EPLB READ did not select UCX") - self._xfer_cache[cache_key] = (local_dlist, remote_dlist, xfer) - self._used_xfer_cache_keys.add(cache_key) - return self._xfer_cache[cache_key] - except Exception: - self._release_xfers([(local_dlist, remote_dlist, xfer)]) - raise - - def _copy_batch(self, batch, prepared_batch) -> None: - if prepared_batch.push_batch is not None: - self._batch_memcpy.enqueue(prepared_batch.push_batch, self._push_stream.cuda_stream) - xfers = [ - self._get_remote_read(src_rank, entries) for src_rank, entries in prepared_batch.remote_entries.items() - ] - self._wait_xfers(xfers) - self._push_stream.synchronize() - # Before a rank publishes this batch it has completed its outgoing source-pushes and - # incoming UCX READs. The manager's global MIN-ready gate therefore means all transfers - # are complete before any rank commits, without a destination-side GPU wait. - - def _start_transfer_generation(self) -> None: - self._used_xfer_cache_keys.clear() - - def _finish_transfer_generation(self) -> None: - errors = [] - for cache_key in set(self._xfer_cache) - self._used_xfer_cache_keys: - xfer = self._xfer_cache[cache_key] - try: - self._release_xfers([xfer]) - except Exception as exc: - unreleased = getattr(exc, "unreleased_xfers", None) - if unreleased: - self._xfer_cache[cache_key] = unreleased[0] - errors.append(exc) - else: - del self._xfer_cache[cache_key] - if errors: - raise RuntimeError("NIXL EPLB cache eviction failed") from errors[0] - - def shutdown(self) -> None: - agent = self._nixl_agent - errors = [] - getattr(self, "_used_xfer_cache_keys", set()).clear() - if agent is not None: - for cache_key, xfer in list(self._xfer_cache.items()): - try: - self._release_xfers([xfer]) - except Exception as exc: - unreleased = getattr(exc, "unreleased_xfers", None) - if unreleased: - self._xfer_cache[cache_key] = unreleased[0] - errors.append(exc) - else: - del self._xfer_cache[cache_key] - if errors: - raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] - for remote_name in list(self._remote_agents.values()): - if agent is not None: - try: - agent.remove_remote_agent(remote_name) - except Exception as exc: - errors.append(exc) - self._remote_agents.clear() - self._remote_layouts.clear() - if agent is not None and self._registered_descs is not None: - try: - agent.deregister_memory(self._registered_descs, backends=["UCX"]) - except Exception as exc: - errors.append(exc) - self._registered_descs = None - self._nixl_agent = None - getattr(self, "_ipc_staging", {}).clear() - if errors: - raise RuntimeError("NIXL EPLB shutdown failed") from errors[0] - - def __del__(self): - try: - self.shutdown() - except Exception: - pass - - -def _commit_staging_rows( - live: torch.Tensor, - staging: torch.Tensor, - num_experts_per_rank: int, - changed_dst_slots: Sequence[int], -) -> None: - slots = sorted(set(changed_dst_slots)) - if not slots: - return - run_start = previous = slots[0] - for dst_slot in (*slots[1:], None): - if dst_slot is not None and dst_slot == previous + 1: - previous = dst_slot - continue - run_length = previous - run_start + 1 - live.narrow(0, num_experts_per_rank + run_start, run_length).copy_( - staging.narrow(0, run_start, run_length), non_blocking=True - ) - if dst_slot is not None: - run_start = previous = dst_slot - - -class _CudaMemLocation(ctypes.Structure): - _fields_ = [("type", ctypes.c_int), ("id", ctypes.c_int)] - - -class _CudaMemcpyAttributes(ctypes.Structure): - _fields_ = [ - ("srcAccessOrder", ctypes.c_int), - ("srcLocHint", _CudaMemLocation), - ("dstLocHint", _CudaMemLocation), - ("flags", ctypes.c_uint), - ] - - -@dataclass -class _PreparedCudaMemcpyBatch: - """Host-side arrays retained for one cudaMemcpyBatchAsync submission.""" - - dsts: object - srcs: object - sizes: object - attrs: _CudaMemcpyAttributes - attrs_idxs: object - count: int - - -class _CudaBatchMemcpy: - """CUDA 13.x ``cudaMemcpyBatchAsync`` binding for EPLB source-push.""" - - _SRC_ACCESS_ORDER_STREAM = 1 - _PREFER_OVERLAP_WITH_COMPUTE = 1 - _CUDA_13_0 = 13000 - _CUDA_14_0 = 14000 - - def __init__(self, library=None): - if library is None: - path = self._find_loaded_cudart() - if path is None: - raise RuntimeError( - "NIXL same-node source-push requires CUDA Runtime 13.x cudaMemcpyBatchAsync; " - "libcudart.so.13 is not loaded" - ) - try: - library = ctypes.CDLL(path) - except OSError as exc: - raise RuntimeError(f"cannot load libcudart: {exc}") from exc - - try: - runtime_get_version = library.cudaRuntimeGetVersion - self._batch_async = library.cudaMemcpyBatchAsync - self._get_error_string = library.cudaGetErrorString - except AttributeError as exc: - raise RuntimeError("cudaMemcpyBatchAsync is unavailable") from exc - - runtime_get_version.restype = ctypes.c_int - runtime_get_version.argtypes = [ctypes.POINTER(ctypes.c_int)] - self._get_error_string.restype = ctypes.c_char_p - self._get_error_string.argtypes = [ctypes.c_int] - runtime_version = ctypes.c_int() - result = runtime_get_version(ctypes.byref(runtime_version)) - if result != 0: - raise RuntimeError(f"cudaRuntimeGetVersion failed with CUDA error {result}") - if not self._CUDA_13_0 <= runtime_version.value < self._CUDA_14_0: - raise RuntimeError( - f"cudaMemcpyBatchAsync requires CUDA Runtime 13.x (13.0 ABI), found {runtime_version.value}" - ) - - pointer_array = ctypes.POINTER(ctypes.c_void_p) - self._batch_async.restype = ctypes.c_int - self._batch_async.argtypes = [ - pointer_array, - pointer_array, - ctypes.POINTER(ctypes.c_size_t), - ctypes.c_size_t, - ctypes.POINTER(_CudaMemcpyAttributes), - ctypes.POINTER(ctypes.c_size_t), - ctypes.c_size_t, - ctypes.c_void_p, - ] - - @staticmethod - def prepare(copies: Iterable[Tuple[int, int, int]]) -> _PreparedCudaMemcpyBatch: - copies = tuple(copies) - if not copies: - raise ValueError("cudaMemcpyBatchAsync requires at least one copy") - for src, dst, size in copies: - if not src or not dst or size <= 0: - raise ValueError("cudaMemcpyBatchAsync requires non-null pointers and positive sizes") - count = len(copies) - dsts = (ctypes.c_void_p * count)(*(dst for _, dst, _ in copies)) - srcs = (ctypes.c_void_p * count)(*(src for src, _, _ in copies)) - sizes = (ctypes.c_size_t * count)(*(size for _, _, size in copies)) - attrs = _CudaMemcpyAttributes() - attrs.srcAccessOrder = _CudaBatchMemcpy._SRC_ACCESS_ORDER_STREAM - attrs.flags = _CudaBatchMemcpy._PREFER_OVERLAP_WITH_COMPUTE - attrs_idxs = (ctypes.c_size_t * 1)(0) - return _PreparedCudaMemcpyBatch(dsts, srcs, sizes, attrs, attrs_idxs, count) - - def enqueue(self, prepared: _PreparedCudaMemcpyBatch, stream: int) -> None: - result = self._batch_async( - prepared.dsts, - prepared.srcs, - prepared.sizes, - prepared.count, - ctypes.byref(prepared.attrs), - prepared.attrs_idxs, - 1, - ctypes.c_void_p(stream), - ) - if result != 0: - message = self._get_error_string(result) - error = message.decode("utf-8") if message else f"CUDA error {result}" - raise RuntimeError(f"cudaMemcpyBatchAsync failed: {error}") - - @staticmethod - def _find_loaded_cudart() -> Optional[str]: - """Return a mapped CUDA 13 runtime without loading CUDA as a side effect.""" - try: - with open("/proc/self/maps") as maps: - for line in maps: - if "libcudart" not in line: - continue - path_start = line.find("/") - if path_start < 0: - continue - path = line[path_start:].strip().removesuffix(" (deleted)") - if re.search(r"libcudart[^/]*\.so\.13(?:\D|$)", path): - return path - except OSError: - pass - return None diff --git a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py index 1e170975a3..125edede25 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/base_kv_move_manager.py @@ -44,12 +44,6 @@ def __init__( start_func=start_trans_process_func, up_status_in_queue=up_status_in_queue, ) - - for trans_process in self.kv_trans_processes: - if not trans_process.wait_until_ready(): - raise RuntimeError(f"KV trans module for device {trans_process.device_id} failed to initialize") - - for trans_process in self.kv_trans_processes: threading.Thread(target=self.task_ret_handle_loop, args=(trans_process,), daemon=True).start() # 通过 io buffer 将命令写入到推理进程中 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py index fcfbaf1936..121a528a43 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_kv_move_manager.py @@ -21,11 +21,8 @@ def start_decode_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - while not event.wait(timeout=1): - if not proc.is_alive(): - raise RuntimeError("decode kv move manager process failed during initialization") - if not proc.is_alive(): - raise RuntimeError("decode kv move manager process exited during initialization") + event.wait() + assert proc.is_alive() logger.info("decode kv move manager process started") return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index 5eeeabb40f..f015264a8b 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -47,7 +47,6 @@ def _init_env( task_out_queue: mp.Queue, up_status_in_queue: Optional[mp.SimpleQueue], ): - module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -101,15 +100,11 @@ def _init_env( up_status_in_queue=up_status_in_queue, ) assert manager is not None - task_out_queue.put("module_ready") - module_ready = True while True: time.sleep(100) except Exception as e: - if not module_ready: - task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass @@ -146,7 +141,6 @@ def __init__( self.recv_task_group_queue = queue.Queue() self.waiting_dict_lock = threading.Lock() self.waiting_dict: Dict[str, PDChunckedTransTask] = {} - self.request_last_progress_time: Dict[int, float] = {} self.request_page_task_queue = queue.Queue() self.ready_page_task_queue = queue.Queue() self.success_queue = queue.Queue() @@ -193,10 +187,10 @@ def _warmup(self): def recv_task_loop(self): while True: obj: Union[PDChunckedTransTaskGroup, PDAbortReq] = self.task_in_queue.get() - if isinstance(obj, (PDChunckedTransTaskGroup, PDAbortReq)): - # Keep the producer's group-before-abort order through dispatch as well. - # Aborting here can miss a group that is still waiting in this queue. + if isinstance(obj, PDChunckedTransTaskGroup): self.recv_task_group_queue.put(obj) + elif isinstance(obj, PDAbortReq): + self._abort(request_id=obj.request_id) else: assert False, f"recv error obj {obj}" @@ -217,10 +211,7 @@ def _abort(self, request_id: int, error_info: str = "aborted req"): @log_exception def dispatch_task_loop(self): while True: - trans_task_group: Union[PDChunckedTransTaskGroup, PDAbortReq] = self.recv_task_group_queue.get() - if isinstance(trans_task_group, PDAbortReq): - self._abort(request_id=trans_task_group.request_id) - continue + trans_task_group: PDChunckedTransTaskGroup = self.recv_task_group_queue.get() with self.waiting_dict_lock: for task in trans_task_group.task_list: @@ -252,19 +243,6 @@ def dispatch_task_loop(self): self.up_status_in_queue.put(up_status) - def _pop_waiting_task_for_notify(self, notify_task: PDChunckedTransTask): - with self.waiting_dict_lock: - local_trans_task = self.waiting_dict.pop(notify_task.get_key(), None) - if local_trans_task is None: - return None - - # Decode creates every page task before prefill starts producing pages. - # A matched notify is forward progress for the request, so future pages - # use an idle timeout instead of their original creation time. - self.request_last_progress_time[local_trans_task.request_id] = time.time() - - return local_trans_task - @log_exception def accept_peer_task_loop( self, @@ -313,7 +291,8 @@ def accept_peer_task_loop( # 到了请求页面的阶段 remote_trans_task = notify_obj if remote_trans_task.write_stage == "request": - local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) + with self.waiting_dict_lock: + local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) if local_trans_task is not None: local_trans_task.prefill_agent_name = remote_trans_task.prefill_agent_name local_trans_task.prefill_agent_metadata = remote_trans_task.prefill_agent_metadata @@ -343,7 +322,8 @@ def accept_peer_task_loop( # prefill 写完数据到了 done 阶段 if remote_trans_task.write_stage == "done": - local_trans_task = self._pop_waiting_task_for_notify(remote_trans_task) + with self.waiting_dict_lock: + local_trans_task = self.waiting_dict.pop(remote_trans_task.get_key(), None) if local_trans_task is not None: local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob @@ -372,23 +352,9 @@ def accept_peer_task_loop( def _check_tasks_time_out(self): with self.waiting_dict_lock: timeout_tasks = [] - pending_request_ids = set() - now = time.time() for key, trans_task in list(self.waiting_dict.items()): - if trans_task.start_trans_time is None: - request_last_progress = self.request_last_progress_time.get(trans_task.request_id) - is_timeout = ( - request_last_progress is not None and now - request_last_progress > trans_task.time_out_secs - ) - else: - is_timeout = trans_task.time_out() - if is_timeout: + if trans_task.time_out(): timeout_tasks.append(self.waiting_dict.pop(key)) - elif trans_task.start_trans_time is None: - pending_request_ids.add(trans_task.request_id) - for request_id in list(self.request_last_progress_time): - if request_id not in pending_request_ids: - self.request_last_progress_time.pop(request_id) for trans_task in timeout_tasks: trans_task.error_info = "time out in accept_peer_task_loop" diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py index e09ba79c95..b23a5c4141 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_kv_move_manager.py @@ -20,11 +20,8 @@ def start_prefill_kv_move_manager_process(args, info_queue: mp.Queue): event = mp.Event() proc = mp.Process(target=_init_env, args=(args, info_queue, event)) proc.start() - while not event.wait(timeout=1): - if not proc.is_alive(): - raise RuntimeError("prefill kv move manager process failed during initialization") - if not proc.is_alive(): - raise RuntimeError("prefill kv move manager process exited during initialization") + event.wait() + assert proc.is_alive() logger.info("prefill kv move manager process started") return diff --git a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py index bd728b47fa..91b575cb40 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/prefill_node_impl/prefill_trans_process.py @@ -41,7 +41,6 @@ def _init_env( task_in_queue: mp.Queue, task_out_queue: mp.Queue, ): - module_ready = False install_fatal_thread_excepthook() start_parent_check_thread() import lightllm.utils.rpyc_fix_utils as _ @@ -75,15 +74,11 @@ def _init_env( mem_managers=mem_managers, ) assert manager is not None - task_out_queue.put("module_ready") - module_ready = True while True: time.sleep(100) except Exception as e: - if not module_ready: - task_out_queue.put("init_failed") logger.exception(str(e)) logger.error(f"Fatal error happened in kv trans process: {e}") pass diff --git a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py index 84173a44ba..073ecf23d2 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/trans_process_obj.py @@ -1,4 +1,3 @@ -import queue import threading import psutil import torch.multiprocessing as mp @@ -59,23 +58,5 @@ def is_trans_process_health(self): except: return False - def wait_until_ready(self): - for _ in range(600): - try: - status = self.task_out_queue.get(timeout=1) - except queue.Empty: - if not self.process.is_alive(): - logger.error(f"KV trans process for device {self.device_id} exited during initialization") - return False - continue - - if status != "module_ready": - logger.error(f"KV trans module for device {self.device_id} failed to initialize: {status}") - return False - return True - - logger.error(f"KV trans module for device {self.device_id} initialization timed out") - return False - def killself(self): self.process.kill() diff --git a/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py b/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py new file mode 100644 index 0000000000..596eca4f24 --- /dev/null +++ b/lightllm/server/router/model_infer/mode_backend/redundancy_expert_manager.py @@ -0,0 +1,158 @@ +# 对于 deepseekv3 模型在 ep 运行模式下,自动分析统计各个专家的出现频率,然后 +# 自动更新当前的冗余专家为新的冗余专家。 +import torch +import time +import enum +import lightllm.utils.petrel_helper as utils +import threading +import json +from typing import List +from lightllm.common.basemodel.basemodel import TpPartBaseModel +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_redundancy import ( + FusedMoeWeightEPAutoRedundancy, +) +from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.fused_moe_weight import FusedMoeWeight +from lightllm.utils.envs_utils import get_env_start_args, get_redundancy_expert_update_interval +from lightllm.utils.envs_utils import get_redundancy_expert_update_max_load_count +from lightllm.utils.envs_utils import get_redundancy_expert_num +from lightllm.utils.dist_utils import get_global_rank +from lightllm.common.basemodel.layer_weights.hf_load_utils import load_func +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +class RedundancyExpertManager: + def __init__(self, model: TpPartBaseModel): + self.args = get_env_start_args() + self.model = model + self.ep_fused_moeweights: List[FusedMoeWeightEPAutoRedundancy] = [] + for layer in self.model.trans_layers_weight: + ep_weights = self._find_members_of_class(layer, FusedMoeWeight) + assert len(ep_weights) <= 1 + self.ep_fused_moeweights.extend([FusedMoeWeightEPAutoRedundancy(e) for e in ep_weights]) + + # save load params + self.use_safetensors = True + files = utils.PetrelHelper.list(self.args.model_dir, extension="all") + candidate_files = list(filter(lambda x: x.endswith(".safetensors"), files)) + if len(candidate_files) == 0: + self.use_safetensors = False + candidate_files = list(filter(lambda x: x.endswith(".bin"), files)) + assert len(candidate_files) != 0, "can only support pytorch tensor and safetensors format for weights." + self.candidate_files = candidate_files + + # state 1. check_to_update 2. prepare_update 3. start_load_hf_weights 4. wait_load_ready, 5. commit + self.state: _STATE = _STATE.CHECK_TO_UPDATE + self.update_time = time.time() + self.update_interval = get_redundancy_expert_update_interval() + self.load_thread: threading.Thread = None + self.global_rank = get_global_rank() + # 冗余专家的最大加载次数 + self.load_count = 0 + self.max_load_count = get_redundancy_expert_update_max_load_count() + + # 清理counter + self._clear_all_counter() + + self.rank0_redundancy_expert_config = { + "redundancy_expert_num": get_redundancy_expert_num(), + "default": list(range(get_redundancy_expert_num())), + } + + def step(self): + if self.load_count >= self.max_load_count: + return + + if self.state == _STATE.CHECK_TO_UPDATE: + cur_time = time.time() + if cur_time - self.update_time > self.update_interval: + self.update_time = cur_time + self.state = _STATE.PREPARE_UPDATE + logger.info(f"global_rank {self.global_rank} state to prepare update") + elif self.state == _STATE.PREPARE_UPDATE: + self._prepare_load_new_redundancy_expert() + self.state = _STATE.START_LOAD_HF_WEIGHTS + logger.info(f"global_rank {self.global_rank} state to start load hf weights") + + elif self.state == _STATE.START_LOAD_HF_WEIGHTS: + self.load_thread = threading.Thread(target=self._load_hf_weights, daemon=True) + self.load_thread.start() + self.state = _STATE.WAIT_LOAD_READY + logger.info(f"global_rank {self.global_rank} state to wait load ready") + + elif self.state == _STATE.WAIT_LOAD_READY: + if not self.load_thread.is_alive(): + self.load_thread = None + self.state = _STATE.COMMIT + logger.info(f"global_rank {self.global_rank} state to commit") + + elif self.state == _STATE.COMMIT: + self._commit() + self.state = _STATE.CHECK_TO_UPDATE + self.load_count += 1 + logger.info(f"global_rank {self.global_rank} state to check to update") + return + + def _prepare_load_new_redundancy_expert(self): + for w in self.ep_fused_moeweights: + topk_redundancy_expert_ids = w.prepare_redundancy_experts() + if self.global_rank == 0: + self.rank0_redundancy_expert_config[str(w._ep_w.layer_num)] = topk_redundancy_expert_ids + + if self.global_rank == 0: + try: + with open("./redundancy_expert_config.json", "w") as f: + json.dump(self.rank0_redundancy_expert_config, f, indent=4) + logger.info( + f"rank {self.global_rank} save redundancy_expert_config.json to ./redundancy_expert_config.json" + ) + except BaseException as e: + logger.exception(str(e)) + logger.error(f"global rank {self.global_rank} save redundancy_expert_config.json failed") + + return + + def _load_hf_weights(self): + start = time.time() + try: + for file in self.candidate_files: + load_func( + file, + use_safetensors=self.use_safetensors, + pre_post_layer=None, + transformer_layer_list=self.ep_fused_moeweights, + weight_dir=self.args.model_dir, + ) + except BaseException as e: + logger.exception(str(e)) + raise e + cost_time = time.time() - start + logger.info(f"global rank {self.global_rank} load redundancy_expert cost time: {cost_time} s") + return + + def _commit(self): + for w in self.ep_fused_moeweights: + w.commit() + return + + def _find_members_of_class(self, obj, cls): + members = [] + for attr in dir(obj): + value = getattr(obj, attr) + if isinstance(value, cls): + members.append(value) + return members + + def _clear_all_counter(self): + for w in self.ep_fused_moeweights: + w.clear_counter() + return + + +class _STATE(enum.Enum): + CHECK_TO_UPDATE = 0 + PREPARE_UPDATE = 1 + START_LOAD_HF_WEIGHTS = 2 + WAIT_LOAD_READY = 3 + COMMIT = 4 diff --git a/lightllm/server/router/model_infer/model_rpc.py b/lightllm/server/router/model_infer/model_rpc.py index 2c5f066779..670a040556 100644 --- a/lightllm/server/router/model_infer/model_rpc.py +++ b/lightllm/server/router/model_infer/model_rpc.py @@ -26,15 +26,12 @@ PDDecodeNode, PDDPForDecodeNode, ) +from lightllm.server.router.model_infer.mode_backend.redundancy_expert_manager import RedundancyExpertManager from lightllm.server.router.model_infer.mode_backend.rl_backend_ops import RlBackendOps -from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( - EPBalanceMonitor, - should_enable_ep_balance_monitor, -) from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.utils.log_utils import init_logger from lightllm.utils.graceful_utils import graceful_registry -from lightllm.utils.process_check import install_fatal_thread_excepthook, start_parent_check_thread +from lightllm.utils.process_check import start_parent_check_thread from lightllm.utils.envs_utils import get_unique_server_name, get_model_infer_recursion_limit from lightllm.utils.torch_memory_saver_utils import MemoryTag from lightllm.server.io_struct import RlOpReq, RlOpRsp @@ -102,10 +99,12 @@ def exposed_init_model(self, kvargs): self.backend.init_model(kvargs) self.rl_backend_ops = RlBackendOps(self.backend) if self.args.enable_rl else None - if should_enable_ep_balance_monitor(self.args): - monitor = EPBalanceMonitor(self.backend.model) - if monitor.enabled: - self.backend.model.ep_balance_monitor = monitor + # only deepseekv3 can support auto_update_redundancy_expert + if self.args.auto_update_redundancy_expert: + self.redundancy_expert_manager = RedundancyExpertManager(self.backend.model) + logger.info("init redundancy_expert_manager") + else: + self.redundancy_expert_manager = None return def exposed_get_max_total_token_num(self): @@ -173,7 +172,6 @@ def _init_env( socket_path, success_event, ): - install_fatal_thread_excepthook() import lightllm.utils.rpyc_fix_utils as _ # spawn 不继承父进程的递归上限;在启动推理线程前为每个 rank 显式设置。 diff --git a/lightllm/server/visualserver/model_infer/model_rpc.py b/lightllm/server/visualserver/model_infer/model_rpc.py index 95209ef6bc..6f8ebed777 100644 --- a/lightllm/server/visualserver/model_infer/model_rpc.py +++ b/lightllm/server/visualserver/model_infer/model_rpc.py @@ -22,7 +22,6 @@ from lightllm.models.qwen3_vl.qwen3_visual import Qwen3VisionTransformerPretrainedModel from lightllm.models.tarsier2.tarsier2_visual import TarsierVisionTransformerPretrainedModel from lightllm.models.qwen3_omni_moe_thinker.qwen3_omni_visual import Qwen3OmniMoeVisionTransformerPretrainedModel -from lightllm.models.neo_chat_moe.neo_visual import NeoVisionTransformerPretrainedModel from lightllm.utils.infer_utils import set_random_seed from lightllm.utils.dist_utils import init_vision_distributed_env from lightllm.utils.envs_utils import get_env_start_args @@ -115,6 +114,8 @@ def exposed_init_model(self, kvargs): .bfloat16() ) elif self.model_type == "neo_chat": + from lightllm.models.neo_chat_moe.neo_visual import NeoVisionTransformerPretrainedModel + self.model = NeoVisionTransformerPretrainedModel(kvargs, **model_cfg["vision_config"]).eval().bfloat16() else: raise Exception(f"can not support {self.model_type} now") diff --git a/lightllm/utils/device_utils.py b/lightllm/utils/device_utils.py index 8cedb001a7..750bdbd4d6 100644 --- a/lightllm/utils/device_utils.py +++ b/lightllm/utils/device_utils.py @@ -45,11 +45,6 @@ def is_sm100_gpu(): return torch.cuda.get_device_capability()[0] == 10 -@lru_cache(maxsize=None) -def is_sm90_gpu(): - return torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 9 - - @lru_cache(maxsize=None) def get_device_sm_regs_num(): import triton diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index 1ac3ae196c..f4e4a41f13 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -150,35 +150,64 @@ def get_lightllm_websocket_max_message_size(): return int(os.getenv("LIGHTLLM_WEBSOCKET_MAX_SIZE", 128 * 1024 * 1024)) +# get_redundancy_expert_ids and get_redundancy_expert_num are primarily +# used to obtain the IDs and number of redundant experts during inference. +# They depend on a configuration file specified by ep_redundancy_expert_config_path, +# which is a JSON formatted text file. +# The content format is as follows: +# { +# "redundancy_expert_num": 1, # Number of redundant experts per rank +# "0": [0], # Key: layer_index (string), +# # Value: list of original expert IDs that are redundant for this layer +# "1": [0], +# "default": [0] # Default list of redundant expert IDs if layer-specific entry is not found +# } + + @lru_cache(maxsize=None) -def get_prefill_eplb_step_interval(): - """Return the number of prefill forwards between EPLB attempts.""" - interval = int(os.getenv("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL", 20)) - if interval <= 0: - raise ValueError("LIGHTLLM_PREFILL_EPLB_STEP_INTERVAL must be greater than 0") - return interval +def get_redundancy_expert_ids(layer_index: int): + """ + Get the redundancy expert ids from the environment variable. + :return: List of redundancy expert ids. + """ + args = get_env_start_args() + if args.ep_redundancy_expert_config_path is None: + return [] + + with open(args.ep_redundancy_expert_config_path, "r") as f: + config = json.load(f) + if str(layer_index) in config: + return config[str(layer_index)] + else: + return config.get("default", []) + + +@lru_cache(maxsize=None) +def get_redundancy_expert_num(): + """ + Get the number of redundancy experts from the environment variable. + :return: Number of redundancy experts. + """ + args = get_env_start_args() + if args.ep_redundancy_expert_config_path is None: + return 0 + + with open(args.ep_redundancy_expert_config_path, "r") as f: + config = json.load(f) + if "redundancy_expert_num" in config: + return config["redundancy_expert_num"] + else: + return 0 @lru_cache(maxsize=None) -def get_eplb_rebalance_gain_threshold() -> float: - """Return the EPLB gain threshold: estimated critical-load reduction ratio; 0.05 means 5%.""" - env_name = "LIGHTLLM_EPLB_REBALANCE_GAIN_THRESHOLD" - raw_value = os.getenv(env_name, "0.05") - value = float(raw_value) - if not 0.0 <= value <= 1.0: - raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") - return value +def get_redundancy_expert_update_interval(): + return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_INTERVAL", 30 * 60)) @lru_cache(maxsize=None) -def get_eplb_placement_stickiness() -> float: - """Return the EPLB placement stickiness: a keep-bonus, as a fraction of the mean per-layer expert load.""" - env_name = "LIGHTLLM_EPLB_PLACEMENT_STICKINESS" - raw_value = os.getenv(env_name, "0.1") - value = float(raw_value) - if not 0.0 <= value <= 1.0: - raise ValueError(f"{env_name} must be a ratio between 0.0 and 1.0, got {raw_value!r}") - return value +def get_redundancy_expert_update_max_load_count(): + return int(os.getenv("LIGHTLLM_REDUNDANCY_EXPERT_UPDATE_MAX_LOAD_COUNT", 1)) @lru_cache(maxsize=None) diff --git a/lightllm/utils/error_utils.py b/lightllm/utils/error_utils.py index 25d7744e50..acf76ad54d 100644 --- a/lightllm/utils/error_utils.py +++ b/lightllm/utils/error_utils.py @@ -10,10 +10,6 @@ class InvalidRequestError(ValueError): """Request validation failed before generation started.""" -class GenerationError(Exception): - """Generation stopped because of an internal server failure.""" - - class ServerBusyError(Exception): """Custom exception for server busy/overload situations""" diff --git a/lightllm/utils/profile_max_tokens.py b/lightllm/utils/profile_max_tokens.py index 73f172d908..5ec2a9145e 100644 --- a/lightllm/utils/profile_max_tokens.py +++ b/lightllm/utils/profile_max_tokens.py @@ -74,14 +74,6 @@ def profile_mtp_weight_memory(model): weight_memory_before = torch.cuda.memory_allocated() yield target_weight_bytes = torch.cuda.memory_allocated() - weight_memory_before - # Draft construction can opt out of portions of the target's weight layout - # which it never instantiates (for example EPLB redundant rows). - excluded_weight_bytes = int(getattr(model, "get_mtp_profile_weight_exclusion", lambda: 0)()) - if not 0 <= excluded_weight_bytes <= target_weight_bytes: - raise ValueError( - f"invalid MTP profile exclusion {excluded_weight_bytes}; measured target weights={target_weight_bytes}" - ) - target_weight_bytes -= excluded_weight_bytes model.mem_fraction = get_mtp_adjusted_mem_fraction( mem_fraction=model.mem_fraction, target_weight_bytes=target_weight_bytes, diff --git a/lightllm/utils/start_utils.py b/lightllm/utils/start_utils.py index f324ddf1f6..3d764d7789 100644 --- a/lightllm/utils/start_utils.py +++ b/lightllm/utils/start_utils.py @@ -1,5 +1,4 @@ import os -import shutil import ctypes import signal import subprocess @@ -19,10 +18,6 @@ class SubmoduleManager: def __init__(self): self.processes = [] self.process_names = {} - self.disk_cache_dir = None - - def register_disk_cache_dir(self, cache_dir): - self.disk_cache_dir = cache_dir def start_submodule_processes(self, start_funcs=[], start_args=[]): assert len(start_funcs) == len(start_args) @@ -30,45 +25,33 @@ def start_submodule_processes(self, start_funcs=[], start_args=[]): processes = [] managed_processes = [] - try: - for start_func, start_arg in zip(start_funcs, start_args): - pipe_reader, pipe_writer = mp.Pipe(duplex=False) - process = mp.Process( - target=start_func, - args=start_arg + (pipe_writer,), - ) - try: - process.start() - finally: - pipe_writer.close() - pipe_readers.append(pipe_reader) - processes.append(process) - managed_process = psutil.Process(process.pid) - managed_processes.append(managed_process) - self.processes.append(managed_process) - self.process_names[managed_process] = managed_process.name() - - # Wait for all processes to initialize - for index, pipe_reader in enumerate(pipe_readers): - try: - init_state = pipe_reader.recv() - finally: - pipe_reader.close() - if init_state != "init ok": - logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") - raise SystemExit(1) - logger.info(f"init func {start_funcs[index].__name__} : {str(init_state)}") - - assert all([proc.is_alive() for proc in processes]) - except BaseException: - for proc in processes: - if proc.is_alive(): + for start_func, start_arg in zip(start_funcs, start_args): + pipe_reader, pipe_writer = mp.Pipe(duplex=False) + process = mp.Process( + target=start_func, + args=start_arg + (pipe_writer,), + ) + process.start() + pipe_readers.append(pipe_reader) + processes.append(process) + # 初始化完成前也可能收到退出信号,因此子进程启动后立即纳入管理。 + managed_process = psutil.Process(process.pid) + managed_processes.append(managed_process) + self.processes.append(managed_process) + self.process_names[managed_process] = managed_process.name() + + # Wait for all processes to initialize + for index, pipe_reader in enumerate(pipe_readers): + init_state = pipe_reader.recv() + if init_state != "init ok": + logger.error(f"init func {start_funcs[index].__name__} : {str(init_state)}") + for proc in processes: proc.kill() - for proc in processes: - proc.join() - self.terminate_all_processes() - raise + sys.exit(1) + else: + logger.info(f"init func {start_funcs[index].__name__} : {str(init_state)}") + assert all([proc.is_alive() for proc in processes]) return managed_processes def register_process_tree(self, root_process): @@ -108,19 +91,9 @@ def terminate_all_processes(self): if alive_pids: logger.warning(f"Processes still alive after SIGKILL: {alive_pids}") - # LightMem owns files under this directory, so remove it only after the cache process has exited. - if self.disk_cache_dir is not None: - try: - shutil.rmtree(self.disk_cache_dir) - except FileNotFoundError: - pass - except Exception as e: - logger.warning(f"Failed to remove disk cache directory {self.disk_cache_dir}: {e}") - else: - logger.info(f"Removed disk cache directory {self.disk_cache_dir}") - # recover the gpu compute mode - if get_env_start_args().enable_mps: + is_enable_mps = get_env_start_args().enable_mps + if is_enable_mps: from lightllm.utils.device_utils import stop_mps stop_mps() @@ -149,11 +122,10 @@ def setup_signal_handlers(self, http_server_process=None): def signal_handler(sig, _frame): if sig == signal.SIGINT: logger.info("Received SIGINT (Ctrl+C), forcing immediate exit...") - try: - if http_server_process is not None: - kill_recursive(http_server_process) - finally: - self.terminate_all_processes() + if http_server_process is not None: + kill_recursive(http_server_process) + + self.terminate_all_processes() logger.info("All processes have been forcefully terminated.") sys.exit(0) @@ -162,17 +134,16 @@ def signal_handler(sig, _frame): else: logger.info("Received SIGHUP (terminal closed), shutting down gracefully...") - try: - if http_server_process is not None and http_server_process.poll() is None: - http_server_process.send_signal(signal.SIGTERM) - try: - http_server_process.wait(timeout=60) - logger.info("HTTP server exited gracefully") - except subprocess.TimeoutExpired: - logger.warning("HTTP server did not exit in time, killing it...") - kill_recursive(http_server_process) - finally: - self.terminate_all_processes() + if http_server_process is not None and http_server_process.poll() is None: + http_server_process.send_signal(signal.SIGTERM) + try: + http_server_process.wait(timeout=60) + logger.info("HTTP server exited gracefully") + except subprocess.TimeoutExpired: + logger.warning("HTTP server did not exit in time, killing it...") + kill_recursive(http_server_process) + + self.terminate_all_processes() logger.info("All processes have been terminated gracefully.") sys.exit(0) diff --git a/revert.md b/revert.md new file mode 100644 index 0000000000..20fd8d2cfe --- /dev/null +++ b/revert.md @@ -0,0 +1,47 @@ +# support_ds4 拆分时排除的提交 + +本分支从 `support_ds4@7369626cc3487cacaaccad611943d61f9558635c` 派生,并用一个汇总清理提交移除不准备合入的内容。清理后的工作树只保留 DeepSeek-V4 模型支持相关能力(基础模型、MTP、DSpark、Vision、PD、多级缓存、DP/DPEP、硬件与量化适配、协议与性能优化)。下列 `support_ds4` 提交属于明确排除的专家能力或通用稳定性/运维改动。 + +## 专家能力 + +| 原提交 | 标题 | 处理 | +| --- | --- | --- | +| `c899b318b29786a39dbb9dd2bc63e186e3779687` | `add triton ep backend` | 排除 Triton EP 后端 | +| `33cd0c9d13f56cee332b533058c19685c923ecc9` | `feat: support EPLB for DeepSeek v4` | 排除 DeepSeek-V4 EPLB | +| `7722b82eb210136bd69cd5116ced688e3b8b6424` | `feat: add EPLB` | 排除通用 EPLB 基础设施 | +| `c8f53bebe1bb1bea92174b4bda6e3f12c34612de` | `support fp8xfp8 mega moe` | 排除 Mega-MoE 扩展 | +| `a3160f25b30bdad7856e0e77b7daff1254de9717` | `feat: imbalance statistics` | 排除专家负载不均衡监控 | +| `1a2f29d05ec2ae90fe608ff69e97995df40d6ba6` | `mega_moe support clamp_limit` | 排除 Mega-MoE clamp 扩展 | + +`origin/main` 已有的 Mega-MoE 基础路径不回退;这里只排除 `support_ds4` 新增的 Mega-MoE 能力。 + +## 通用稳定性与运维 + +| 原提交 | 标题 | +| --- | --- | +| `9bb68ebabfb26a0579b8764155bcb5ddcd26380b` | `fix abort` | +| `9a3f99dcf3c492797105362d229479e29db71211` | `fix(pd): clean up requests correctly after node disconnect` | +| `33da96f290e6d58aba60feb658c8aecd84d61384` | `fix(pd): harden control sends and reconnect cleanup` | +| `2a7e6bd0f4ec6d2edda947c43c4d2d032ea97e09` | `PD nodes unavailable: 503` | +| `3d910fe7a6001f1494b9cb52f880ff68d6ee4db5` | `avoid premature KV transfer aborts and surface generation failures` | +| `ef0df4317f6e9db3e19c3366c8e2b2624bc5fb73` | `fix(pd): wait for KV transfer modules to become ready` | +| `aa9b2a34dc9ff310861dc586eae215ebb64dda69` | `fix(api): return 400 for oversized guided JSON schemas` | +| `8ed3da6e5c36ee4c7a722cce11c1eeb0f1d08f3d` | `fix(pd): measure decode KV timeout from request progress` | +| `7bc0955b20816d407748dcc927f461661609c155` | `fix(pd): handle deep prompt-cache keys iteratively` | +| `7229e567526c9428070e03bd025c138137d4390a` | `clean disk cache on shutdown` | +| `811230e0755f2922cf18c63dcfedfe75453da758` | `fix ZMQ crash when cache workers finish concurrently` | +| `7fad48a98421ccb326d529029bc5fa496c3512ce` | `fix(pd): prevent PD master hang on prefill timeout race` | +| `5e491fb5db02d01a5476dad4a35be9f722d7c1fb` | `feat(metrics): add PD pipeline latency metrics` | +| `095c043d4781c8cc8f8b515d4e4e618958ff5a6f` | `perf(openai): reduce chat streaming serialization overhead` | +| `7c65cd00ac66f0998f64709945de18c7139bb3d2` | `fix(pd): preserve aborts received before request registration` | + +若同类修复已经存在于 `origin/main`,本次只去掉 `support_ds4` 的额外增量,不回退主分支已有行为。 + +## 混合提交的拆分处理 + +| 原提交 | 保留 | 排除 | +| --- | --- | --- | +| `baadd7ea316861c089956d46021e0b216a51bcf8` (`perf(server): reduce PD streaming and parser overhead`) | DeepSeek-V4 DSML 流式解析所需逻辑 | 通用 PD streaming/parser 性能改动 | +| `ca4f075860e00649ea73554780b33a417ead21bc` (`synchronize PDL top-k and fail fast on model thread errors`) | DeepSeek-V4 PDL top-k 同步 | 通用 fatal thread excepthook | + +整理阶段产生的临时 revert 提交已经压平,因此最终分支相对 `7369626` 只多一个汇总清理提交;本文件记录被排除的原始提交,便于后续追溯。 diff --git a/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json b/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json new file mode 100644 index 0000000000..241ab25ea3 --- /dev/null +++ b/test/advanced_config/redundancy_expert/test_redundancy_expert_config.json @@ -0,0 +1,180 @@ +{ + "redundancy_expert_num": 1, + "default": [ + 0 + ], + "3": [ + 226 + ], + "4": [ + 123 + ], + "5": [ + 187 + ], + "6": [ + 138 + ], + "7": [ + 132 + ], + "8": [ + 240 + ], + "9": [ + 4 + ], + "10": [ + 88 + ], + "11": [ + 60 + ], + "12": [ + 161 + ], + "13": [ + 178 + ], + "14": [ + 80 + ], + "15": [ + 144 + ], + "16": [ + 195 + ], + "17": [ + 251 + ], + "18": [ + 226 + ], + "19": [ + 87 + ], + "20": [ + 149 + ], + "21": [ + 45 + ], + "22": [ + 214 + ], + "23": [ + 41 + ], + "24": [ + 46 + ], + "25": [ + 156 + ], + "26": [ + 112 + ], + "27": [ + 185 + ], + "28": [ + 58 + ], + "29": [ + 156 + ], + "30": [ + 147 + ], + "31": [ + 199 + ], + "32": [ + 16 + ], + "33": [ + 188 + ], + "34": [ + 227 + ], + "35": [ + 136 + ], + "36": [ + 84 + ], + "37": [ + 15 + ], + "38": [ + 204 + ], + "39": [ + 96 + ], + "40": [ + 226 + ], + "41": [ + 25 + ], + "42": [ + 69 + ], + "43": [ + 122 + ], + "44": [ + 152 + ], + "45": [ + 113 + ], + "46": [ + 98 + ], + "47": [ + 68 + ], + "48": [ + 13 + ], + "49": [ + 102 + ], + "50": [ + 214 + ], + "51": [ + 201 + ], + "52": [ + 182 + ], + "53": [ + 235 + ], + "54": [ + 162 + ], + "55": [ + 125 + ], + "56": [ + 62 + ], + "57": [ + 121 + ], + "58": [ + 105 + ], + "59": [ + 236 + ], + "60": [ + 117 + ] +} \ No newline at end of file diff --git a/test/test_api/test_server_busy_handling.py b/test/test_api/test_server_busy_handling.py index d524fe3c53..cc0181961f 100644 --- a/test/test_api/test_server_busy_handling.py +++ b/test/test_api/test_server_busy_handling.py @@ -66,12 +66,7 @@ def test_stream_starts_response_after_first_chunk_by_default(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), - ) - monkeypatch.setattr( - api_stream_obj, - "_record_pd_send_metrics", - lambda *_: pytest.fail("normal-mode streams must not emit PD-master metrics"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) async def run(): @@ -106,7 +101,7 @@ def test_disable_delay_response_start_sends_status_immediately(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=True, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=True), ) async def run(): @@ -135,7 +130,7 @@ def test_stream_propagates_error_before_response_start(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) async def run(): @@ -161,7 +156,7 @@ def test_delayed_stream_can_return_http_429(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) app = FastAPI() @@ -195,7 +190,7 @@ def test_delayed_stream_can_return_http_400_for_invalid_request(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) app = FastAPI() app.exception_handler(InvalidRequestError)(api_http.invalid_request_exception_handler) @@ -227,7 +222,7 @@ def test_pd_master_anthropic_stream_preserves_error_envelope(monkeypatch): monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) async def anthropic_messages_impl(_request): @@ -328,7 +323,7 @@ def test_delayed_stream_returns_http_400_for_value_error_before_first_chunk(monk monkeypatch.setattr( api_stream_obj, "get_env_start_args", - lambda: SimpleNamespace(disable_delay_response_start=False, run_mode="normal"), + lambda: SimpleNamespace(disable_delay_response_start=False), ) app = FastAPI() app.exception_handler(InvalidRequestError)(api_http.invalid_request_exception_handler) diff --git a/test/test_pd_selector/test_pd_master_multi_choice.py b/test/test_pd_selector/test_pd_master_multi_choice.py index d4999b047a..38f43232fa 100644 --- a/test/test_pd_selector/test_pd_master_multi_choice.py +++ b/test/test_pd_selector/test_pd_master_multi_choice.py @@ -77,7 +77,7 @@ async def wait_to_token_package( manager._wait_to_token_package = wait_to_token_package - with patch.object(HttpServerManager, "_check_and_repair_length", new=MagicMock()): + with patch.object(HttpServerManager, "_check_and_repair_length", new=AsyncMock()): results = [] async for result in manager._generate("prompt", sampling_params, multimodal_params, request): results.append(result) @@ -142,7 +142,7 @@ async def generate_one( manager._generate_one = generate_one - with patch.object(HttpServerManager, "_check_and_repair_length", new=MagicMock()): + with patch.object(HttpServerManager, "_check_and_repair_length", new=AsyncMock()): results = [] async for result in manager._generate("prompt", sampling_params, multimodal_params, request): results.append(result) @@ -180,7 +180,7 @@ async def generate_one(*_args, **_kwargs): manager._generate_one = generate_one - with patch.object(HttpServerManager, "_check_and_repair_length", new=MagicMock()): + with patch.object(HttpServerManager, "_check_and_repair_length", new=AsyncMock()): results = [] async for result in manager._generate("prompt", sampling_params, multimodal_params, request): results.append(result) diff --git a/unit_tests/common/basemodel/test_cuda_graph_layout.py b/unit_tests/common/basemodel/test_cuda_graph_layout.py index 30d4c54368..c1db92bf39 100644 --- a/unit_tests/common/basemodel/test_cuda_graph_layout.py +++ b/unit_tests/common/basemodel/test_cuda_graph_layout.py @@ -2,12 +2,14 @@ from types import SimpleNamespace import pytest +import torch 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.infer_struct import InferStateInfo from lightllm.common.basemodel.mtp_manager import MtpManager @@ -83,6 +85,46 @@ def test_batch_step_size_after_split_controls_capture_range(_graph_args): ] +def test_token_forward_keeps_mtp_hidden_capture_buffer_for_replay(): + model = TpPartBaseModel.__new__(TpPartBaseModel) + model.layers_num = 0 + model.layers_infer = [] + model.trans_layers_weight = [] + model.pre_post_weight = None + model.pre_infer = SimpleNamespace( + token_forward=lambda input_ids, infer_state, layer_weight: infer_state.mtp_draft_input_hiddens, + _tpsp_sp_split=lambda input, infer_state: input, + ) + model.post_infer = SimpleNamespace( + _tpsp_allgather=lambda input, infer_state: input, + token_forward=lambda input, infer_state, layer_weight: object(), + ) + output = SimpleNamespace(to_no_ref_tensor=lambda: None) + model._create_model_output = lambda post_output, infer_state: output + + def make_state(hidden, is_cuda_graph): + state = InferStateInfo() + state.input_ids = torch.zeros(hidden.shape[0], dtype=torch.int64) + state.mtp_draft_input_hiddens = hidden + state.is_cuda_graph = is_cuda_graph + state.hidden_collector = SimpleNamespace(add_final_hidden=lambda value: None) + state.decode_att_state = SimpleNamespace(copy_for_decode_cuda_graph=lambda value: None) + return state + + capture_hidden = torch.zeros((2, 4)) + graph_state = make_state(capture_hidden, is_cuda_graph=True) + model._token_forward(graph_state) + + assert graph_state.mtp_draft_input_hiddens is capture_hidden + replay_state = make_state(torch.ones((2, 4)), is_cuda_graph=False) + graph_state.copy_for_cuda_graph(replay_state) + assert graph_state.mtp_draft_input_hiddens.sum().item() == 8 + + eager_state = make_state(torch.ones((2, 4)), is_cuda_graph=False) + model._token_forward(eager_state) + assert eager_state.mtp_draft_input_hiddens is None + + @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)]) diff --git a/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py b/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py new file mode 100644 index 0000000000..16131ef935 --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_redundancy_topk_ids_repair.py @@ -0,0 +1,151 @@ +import torch +import pytest +from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import redundancy_topk_ids_repair +from lightllm.common.basemodel.triton_kernel.redundancy_topk_ids_repair import expert_id_counter +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +def test_redundancy_topk_ids_repair(): + ep_expert_num = 4 + global_rank = 0 + redundancy_expert_num = 1 + topk_ids = torch.tensor( + [ + [0, 1, 2, 3], + [7, 9, 10, 11], + [1, 3, 5, 7], + ], + dtype=torch.int64, + device="cuda", + ) + + redundancy_expert_ids = torch.tensor( + [ + 0, + ], + dtype=torch.int64, + device="cuda", + ) + + expert_id_counter = torch.zeros(12, dtype=torch.int64, device="cuda") + + redundancy_topk_ids_repair( + topk_ids=topk_ids, + redundancy_expert_ids=redundancy_expert_ids, + ep_expert_num=ep_expert_num, + global_rank=global_rank, + expert_counter=expert_id_counter, + enable_counter=True, + ) + + ans_topk_ids = torch.tensor( + [ + [0, 1, 2, 3], + [7, 9, 10, 11], + [1, 3, 5, 7], + ], + dtype=torch.int64, + device="cuda", + ) + ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids + new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids + ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( + (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 + ) + + assert torch.equal(topk_ids, ans_topk_ids) + assert torch.equal( + expert_id_counter, torch.tensor([1, 2, 1, 2, 0, 1, 0, 2, 0, 1, 1, 1], dtype=torch.int64, device="cuda") + ) + + ep_expert_num = 4 + global_rank = 1 + redundancy_expert_num = 1 + topk_ids = torch.tensor( + [ + [0, 1, 2, 3], + [7, 9, 10, 11], + [1, 3, 5, 7], + ], + dtype=torch.int64, + device="cuda", + ) + + redundancy_expert_ids = torch.tensor( + [ + 5, + ], + dtype=torch.int64, + device="cuda", + ) + redundancy_topk_ids_repair( + topk_ids=topk_ids, + redundancy_expert_ids=redundancy_expert_ids, + ep_expert_num=ep_expert_num, + global_rank=global_rank, + ) + + ans_topk_ids = torch.tensor( + [ + [0, 1, 2, 3], + [7, 9, 10, 11], + [1, 3, 5, 7], + ], + dtype=torch.int64, + device="cuda", + ) + ans_topk_ids = (ans_topk_ids // ep_expert_num) * redundancy_expert_num + ans_topk_ids + new_redundancy_expert_ids = (redundancy_expert_ids // ep_expert_num) * redundancy_expert_num + redundancy_expert_ids + ans_topk_ids[ans_topk_ids == new_redundancy_expert_ids[0]] = ( + (ep_expert_num + redundancy_expert_num) * global_rank + ep_expert_num + 0 + ) + + assert torch.equal(topk_ids, ans_topk_ids) + + +def test_expert_id_counter(): + token_num = 256 + tok_ids = torch.randint( + low=0, + high=12, + size=(token_num, 8), + dtype=torch.int64, + device="cuda", + ) + expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") + expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) + + ans_expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") + ids, counts = torch.unique(tok_ids.view(-1), return_counts=True) + ans_expert_counter[ids] = counts + + assert torch.equal(expert_counter, ans_expert_counter) + + # test speed + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for _ in range(100): + tok_ids = torch.randint( + low=0, + high=12, + size=(token_num, 8), + dtype=torch.int64, + device="cuda", + ) + expert_counter = torch.zeros(12, dtype=torch.int64, device="cuda") + expert_id_counter(topk_ids=tok_ids, expert_counter=expert_counter) + graph.replay() + + start_event = torch.cuda.Event(enable_timing=True) + start_event.record() + graph.replay() + end_event = torch.cuda.Event(enable_timing=True) + end_event.record() + torch.cuda.synchronize() + logger.info(f"expert_id_counter time cost: {start_event.elapsed_time(end_event)} ms") + + +if __name__ == "__main__": + pytest.main() diff --git a/unit_tests/common/fused_moe/test_eplb.py b/unit_tests/common/fused_moe/test_eplb.py deleted file mode 100644 index fd656e07c3..0000000000 --- a/unit_tests/common/fused_moe/test_eplb.py +++ /dev/null @@ -1,3665 +0,0 @@ -import builtins -import io -import threading -import time -from collections import deque -from contextlib import contextmanager -from types import SimpleNamespace - -import pytest -import torch - -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, - build_logical_to_physical_map, - build_logical_to_physical_maps_for_layers, - _estimate_rank_load, - plan_redundant_experts, - select_improving_placements, -) -from lightllm.server.api_cli import make_argument_parser -from lightllm.server.core.objs.start_args_type import StartArgs -from lightllm.server.router.model_infer.infer_batch import g_infer_context -from lightllm.server.router.model_infer.mode_backend import ( - eplb_manager as manager_module, -) -from lightllm.server.router.model_infer.mode_backend import ( - eplb_transfer as transfer_module, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( - deepgemm_impl as deepgemm_module, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl import ( - create_fuse_moe_impl, - FuseMoeMarlin, - FuseMoeTriton, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.impl.base_impl import ( - FuseMoeBaseImpl, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe import ( - fused_moe_weight as fused_weight_module, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - disable_eplb_model_init, - is_eplb_model_init_disabled, -) -from lightllm.common.eplb_utils import extract_eplb_expert_tensors -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( - TransferStep, - _CudaBatchMemcpy, - _commit_staging_rows, - align_target_placement, - build_transfer_plan, -) - - -def _test_parallel_state( - *, - eplb=False, - num_logical_experts=128, - world_size=16, - num_redundant_experts_per_rank=1, - route_counter=None, - recording=False, - recorded_sample_count=0, -): - eplb_state = None - if eplb: - if route_counter is None: - route_counter = torch.zeros((2, num_logical_experts), dtype=torch.int64) - initial_layout_world_size = max(world_size, 2) - eplb_state = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( - num_logical_experts, - initial_layout_world_size, - num_redundant_experts_per_rank, - ), - logical_to_physical_map=torch.zeros((num_logical_experts, 2), dtype=torch.int32), - logical_replica_count=torch.ones(num_logical_experts, dtype=torch.int32), - route_counter=route_counter, - recording=recording, - recorded_sample_count=recorded_sample_count, - ) - return ExpertParallelState( - num_logical_experts=num_logical_experts, - world_size=world_size, - eplb=eplb_state, - ) - - -def _validated_expert_parallel_state( - *, - eplb=True, - n_routed_experts=4, - world_size=2, - num_redundant_experts_per_rank=1, - device="cpu", -): - runtime = None - if eplb: - runtime = EPLBState( - num_redundant_experts_per_rank=num_redundant_experts_per_rank, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids( - n_routed_experts, - world_size, - num_redundant_experts_per_rank, - ), - logical_to_physical_map=torch.empty((n_routed_experts, 2), dtype=torch.int32, device=device), - logical_replica_count=torch.ones(n_routed_experts, dtype=torch.int32, device=device), - route_counter=torch.zeros((2, n_routed_experts), dtype=torch.int64, device=device), - ) - return ExpertParallelState( - num_logical_experts=n_routed_experts, - world_size=world_size, - eplb=runtime, - ) - - -def _set_expert_parallel_state(impl, state): - impl.expert_parallel_state = state - impl.eplb = state.eplb - impl._primary_weight_pack_cache = {} - - -def _manual_runtime_rank_load(source_load, placement, node_world_size, alignment): - """Reference the committed runtime logical-to-physical maps on CPU.""" - samples, layers, nodes, num_logical_experts = source_load.shape - ranks, redundant = placement.shape[1:] - num_experts_per_rank = num_logical_experts // ranks - num_physical_experts_per_rank = num_experts_per_rank + redundant - raw = torch.zeros((samples, layers, ranks, num_logical_experts), dtype=torch.float64) - for layer in range(layers): - for source_node in range(nodes): - logical_to_physical, replica_count = build_logical_to_physical_map( - placement[layer], - num_logical_experts, - source_rank=source_node * node_world_size, - node_world_size=node_world_size, - ) - for expert in range(num_logical_experts): - count = int(replica_count[expert].item()) - for physical_id in logical_to_physical[expert, :count].tolist(): - rank = physical_id // num_physical_experts_per_rank - raw[:, layer, rank, expert] += source_load[:, layer, source_node, expert] / count - return (torch.ceil(raw / alignment) * alignment).sum(dim=3) - - -def test_base_call_template_forwards_selection_and_capture_callback(): - class Impl(FuseMoeBaseImpl): - def _select_experts( - self, - input_tensor, - router_logits, - correction_bias, - top_k, - renormalize, - use_grouped_topk, - topk_group, - num_expert_group, - scoring_func, - per_expert_scale=None, - shared_expert_gate=None, - is_prefill=None, - preserve_logical_ids=False, - ): - seen["select"] = {"preserve_logical_ids": preserve_logical_ids} - return "weights", "physical_ids", "logical_ids" - - def _fused_experts( - self, - input_tensor, - w13, - w2, - topk_weights, - topk_ids, - router_logits=None, - is_prefill=None, - ): - seen["fused"] = {"topk_ids": topk_ids} - return "output" - - seen, captured = {}, [] - impl = Impl(4, 0, 1.0, SimpleNamespace()) - result = impl( - "input", - "logits", - "w13", - "w2", - None, - "softmax", - 2, - False, - False, - 0, - 0, - moe_capture_callback=captured.append, - ) - assert result == "output" - assert captured == ["logical_ids"] - assert seen["select"]["preserve_logical_ids"] - assert seen["fused"]["topk_ids"] == "physical_ids" - - -def test_parallel_state_derives_expert_layout(): - state = _validated_expert_parallel_state(eplb=True) - assert state.num_primary_experts_per_rank == 2 - assert state.num_total_physical_experts == 6 - - -def test_factory_selects_all_paths_and_requires_ep_state(monkeypatch): - plain_quant = SimpleNamespace(method_name="none") - marlin_quant = SimpleNamespace(method_name="awq_marlin") - monkeypatch.setattr(FuseMoeMarlin, "create_workspace", lambda self: None) - state = _validated_expert_parallel_state(eplb=False) - ep_impl = create_fuse_moe_impl( - n_routed_experts=4, - num_fused_shared_experts=0, - routed_scaling_factor=1.0, - quant_method=plain_quant, - expert_parallel_state=state, - ) - assert isinstance(ep_impl, deepgemm_module.FuseMoeDeepGEMM) - assert ep_impl.expert_parallel_state is state - assert state.eplb is None - assert isinstance( - create_fuse_moe_impl( - n_routed_experts=4, - num_fused_shared_experts=0, - routed_scaling_factor=1.0, - quant_method=plain_quant, - ), - FuseMoeTriton, - ) - assert isinstance( - create_fuse_moe_impl( - n_routed_experts=4, - num_fused_shared_experts=0, - routed_scaling_factor=1.0, - quant_method=marlin_quant, - ), - FuseMoeMarlin, - ) - - -def test_find_fused_moe_weights_discovers_direct_layer_attributes(monkeypatch): - class FakeFusedMoeWeight: - def __init__(self, layer_num, enable_ep_moe=True): - self.layer_num_ = layer_num - self.enable_ep_moe = enable_ep_moe - - monkeypatch.setattr(manager_module, "FusedMoeWeight", FakeFusedMoeWeight) - first = FakeFusedMoeWeight(3) - alternate = FakeFusedMoeWeight(1) - aliased = FakeFusedMoeWeight(2) - disabled = FakeFusedMoeWeight(0, enable_ep_moe=False) - model = SimpleNamespace( - trans_layers_weight=[ - SimpleNamespace(moe_weight=first), - SimpleNamespace(alternate_moe_weight=alternate), - SimpleNamespace(moe_weight=aliased, alternate_moe_weight=aliased), - SimpleNamespace(moe_weight=disabled), - ] - ) - - assert manager_module._find_fused_moe_weights(model) == [alternate, aliased, first] - - -def test_eplb_redundant_experts_defaults_per_ep_rank(): - parser = make_argument_parser() - - assert parser.parse_args([]).eplb_num_redundant_experts_per_rank == 2 - assert parser.parse_args(["--eplb_num_redundant_experts_per_rank", "3"]).eplb_num_redundant_experts_per_rank == 3 - assert StartArgs().eplb_num_redundant_experts_per_rank == 2 - - -@pytest.mark.parametrize( - ("num_logical_experts", "num_ranks", "num_redundant_experts_per_rank", "expected"), - [ - (8, 4, 2, [[2, 3], [4, 5], [6, 7], [0, 1]]), - (6, 3, 4, [[2, 3, 4, 5], [4, 5, 0, 1], [0, 1, 2, 3]]), - ], -) -def test_build_initial_redundant_expert_ids( - num_logical_experts, - num_ranks, - num_redundant_experts_per_rank, - expected, -): - actual = build_initial_redundant_expert_ids( - num_logical_experts, - num_ranks, - num_redundant_experts_per_rank, - ) - - assert actual.dtype == torch.int64 - assert actual.shape == (num_ranks, num_redundant_experts_per_rank) - assert torch.equal(actual, torch.tensor(expected, dtype=torch.int64)) - - -def test_plan_redundant_experts_never_uses_owner_or_duplicate_rank(): - expert_load = ( - torch.tensor( - [ - [100, 90, 80, 70, 60, 50, 40, 30], - [30, 40, 50, 60, 70, 80, 90, 100], - ] - ) - .unsqueeze(0) - .unsqueeze(2) - ) - placement = plan_redundant_experts(expert_load, num_ranks=4, num_redundant_experts_per_rank=2) - - for layer_placement in placement: - for rank, expert_ids in enumerate(layer_placement.tolist()): - assert len(expert_ids) == len(set(expert_ids)) - assert all(expert_id // 2 != rank for expert_id in expert_ids) - - -def test_plan_redundant_experts_minimizes_samplewise_aligned_critical_load(): - samples = torch.tensor([[[300, 20, 20, 200]], [[100, 300, 40, 160]]]).unsqueeze(2) - placement = plan_redundant_experts(samples, num_ranks=2, num_redundant_experts_per_rank=1, expert_alignment=128) - candidates = [torch.tensor([[[left], [right]]]) for left in (2, 3) for right in (0, 1)] - - def critical(candidate): - return _estimate_rank_load(samples, candidate, expert_alignment=128).max(dim=2).values.sum() - - assert torch.equal(placement, torch.tensor([[[3], [0]]])) - assert critical(placement) == min(critical(candidate) for candidate in candidates) - - -def test_select_improving_placements_rejects_regressing_layer(): - expert_load = torch.tensor([[8649, 5740, 5002, 3441]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - regressing_candidate = torch.tensor([[[1], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, regressing_candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, regressing_candidate).max() - / _estimate_rank_load(expert_load, regressing_candidate).mean() - ) - assert current_ratio.item() == pytest.approx(1.1007, abs=1e-4) - assert candidate_ratio.item() == pytest.approx(1.1184, abs=1e-4) - assert not improved.item() - assert torch.equal(selected, current) - - -def test_select_improving_placements_rejects_near_balance_when_gain_is_below_threshold(): - expert_load = torch.tensor([[1, 2, 1, 17]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[3], [0]]]) - candidate = torch.tensor([[[3], [1]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - assert current_ratio.item() == pytest.approx(1.047619, abs=1e-6) - assert candidate_ratio.item() == pytest.approx(1.0) - assert not improved.item() - assert torch.equal(selected, current) - - -def test_select_improving_placements_accepts_alignment_aware_gain_even_when_current_ranks_are_balanced(): - expert_load = torch.tensor([[100, 129, 100, 129]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [1]]]) - - selected, improved, metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) - - assert metrics["model_imbalance_ratio"] == pytest.approx(1.0) - assert metrics["candidate_rebalance_gain"] == pytest.approx(0.25) - assert improved.item() - assert torch.equal(selected, candidate) - - -def test_select_improving_placements_rejects_insufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 6, 7]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - relative_improvement = (current_ratio - candidate_ratio) / current_ratio - assert current_ratio.item() == pytest.approx(1.4) - assert candidate_ratio.item() == pytest.approx(1.333333, abs=1e-6) - assert relative_improvement.item() == pytest.approx(0.047619, abs=1e-6) - assert not improved.item() - assert torch.equal(selected, current) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=0.04, - ) - - assert improved.item() - assert torch.equal(selected, candidate) - - -@pytest.mark.parametrize("rebalance_gain_threshold", [-0.01, 1.01, float("nan"), float("inf")]) -def test_select_improving_placements_rejects_invalid_rebalance_gain_threshold( - rebalance_gain_threshold, -): - expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - - with pytest.raises(ValueError, match="rebalance_gain_threshold"): - select_improving_placements( - expert_load, - current, - candidate, - rebalance_gain_threshold=rebalance_gain_threshold, - ) - - -def test_select_improving_placements_accepts_sufficient_rebalance_gain(): - expert_load = torch.tensor([[1, 1, 1, 2]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[3], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - current_ratio = _estimate_rank_load(expert_load, current).max() / _estimate_rank_load(expert_load, current).mean() - candidate_ratio = ( - _estimate_rank_load(expert_load, candidate).max() / _estimate_rank_load(expert_load, candidate).mean() - ) - relative_improvement = (current_ratio - candidate_ratio) / current_ratio - assert current_ratio.item() == pytest.approx(1.2) - assert candidate_ratio.item() == pytest.approx(1.0) - assert relative_improvement.item() == pytest.approx(1 / 6) - assert improved.item() - assert torch.equal(selected, candidate) - - -def test_select_improving_placements_rejects_raw_improvement_that_does_not_improve_aligned_compute(): - expert_load = torch.tensor([[1, 1, 1, 8]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - raw_improving_candidate = torch.tensor([[[3], [0]]]) - - _, raw_improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, raw_improving_candidate, rebalance_gain_threshold=0.05 - ) - selected, aligned_improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, - current, - raw_improving_candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) - - assert raw_improved.item() - assert not aligned_improved.item() - assert torch.equal(selected, current) - - -def test_estimate_rank_load_aligns_each_sample_before_accumulation(): - samples = torch.tensor([[[20, 0, 0, 0]], [[20, 0, 0, 0]]]).unsqueeze(2) - placement = torch.tensor([[[2], [0]]]) - - per_sample = _estimate_rank_load(samples, placement, expert_alignment=128) - accumulated = _estimate_rank_load(samples.sum(dim=0, keepdim=True), placement, expert_alignment=128) - - assert torch.equal(per_sample[:, 0], torch.tensor([[128.0, 128.0], [128.0, 128.0]])) - assert torch.equal(per_sample.sum(dim=0)[0], torch.tensor([256.0, 256.0])) - assert torch.equal(accumulated[0, 0], torch.tensor([128.0, 128.0])) - - -def test_select_improving_placements_rejects_lower_ratio_when_critical_is_unchanged(): - samples = torch.tensor( - [ - [[255, 220, 226, 254]], - [[172, 278, 51, 238]], - [[249, 291, 284, 183]], - ] - ).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - mean_inflating_candidate = torch.tensor([[[2], [1]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - samples, - current, - mean_inflating_candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - ) - current_load = _estimate_rank_load(samples, current, expert_alignment=128) - candidate_load = _estimate_rank_load(samples, mean_inflating_candidate, expert_alignment=128) - current_critical = current_load.max(dim=2).values.sum() - candidate_critical = candidate_load.max(dim=2).values.sum() - - assert candidate_load.mean(dim=2).sum() > current_load.mean(dim=2).sum() - assert current_critical == candidate_critical - assert not improved.item() - assert torch.equal(selected, current) - - -def test_select_improving_placements_accepts_five_percent_critical_reduction(): - samples = torch.tensor( - [ - [[13, 352, 348, 141]], - [[287, 175, 236, 179]], - [[316, 99, 266, 353]], - ] - ).unsqueeze(2) - current = torch.tensor([[[2], [0]]]) - candidate = torch.tensor([[[2], [1]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - samples, current, candidate, rebalance_gain_threshold=0.05, expert_alignment=128 - ) - - assert improved.item() - assert torch.equal(selected, candidate) - - -def test_select_improving_placements_rejects_single_layer_gain_below_model_threshold(): - # Layer 0 becomes better, but layer 1 dominates model critical load. The - # aggregate estimated critical-load reduction gain is below 5%, so neither layer may be changed. - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 10]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]], [[2], [0]]]) - candidate = torch.tensor([[[3], [0]], [[2], [0]]]) - - selected, improved, metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - assert not torch.any(improved) - assert torch.equal(selected, current) - assert metrics["candidate_rebalance_gain"] == pytest.approx(0.5 / 13) - assert metrics["candidate_changed_layer_count"] == 1 - - -def test_select_improving_placements_accepts_only_when_model_gain_reaches_threshold(): - expert_load = torch.tensor([[1, 1, 1, 2], [0, 0, 0, 5]]).unsqueeze(0).unsqueeze(2) - current = torch.tensor([[[2], [0]], [[2], [0]]]) - candidate = torch.tensor([[[3], [0]], [[2], [0]]]) - - selected, improved, _metrics, _before_load, _after_load = select_improving_placements( - expert_load, current, candidate, rebalance_gain_threshold=0.05 - ) - - assert torch.equal(improved, torch.tensor([True, False])) - assert torch.equal(selected, candidate) - - -def test_logical_to_physical_map_has_at_most_one_slot_per_rank(): - redundant_expert_ids = torch.tensor([[2, 3], [0, 1]]) - logical_to_physical, replica_count = build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=4) - - assert logical_to_physical.shape == (4, 2) - assert torch.equal(replica_count, torch.full((4,), 2, dtype=torch.int64)) - - -def test_logical_to_physical_map_requires_expert_count_divisible_by_rank_count_without_source_rank(): - redundant_expert_ids = torch.tensor([[0], [1]]) - - with pytest.raises(AssertionError): - build_logical_to_physical_map(redundant_expert_ids, num_logical_experts=5) - - -def test_logical_to_physical_map_prefers_source_node_replicas(): - # Four ranks, two ranks per node, two primary experts/rank and one - # redundant slot/rank. Expert 0 is primary on rank 0 and replicated on - # rank 2 (the other node). - redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) - rank0_map, rank0_count = build_logical_to_physical_map( - redundant, num_logical_experts=8, source_rank=0, node_world_size=2 - ) - rank1_map, rank1_count = build_logical_to_physical_map( - redundant, num_logical_experts=8, source_rank=1, node_world_size=2 - ) - fallback_redundant = torch.tensor([[4], [5], [0], [1], [2], [3]], dtype=torch.int64) - rank4_map, rank4_count = build_logical_to_physical_map(fallback_redundant, 12, source_rank=4, node_world_size=2) - - assert rank0_count[0].item() == rank1_count[0].item() == 1 - assert torch.equal(rank0_map[0, :1], torch.tensor([0])) - assert torch.equal(rank1_map[0, :1], torch.tensor([0])) - assert rank4_count[0].item() == 2 - assert set(rank4_map[0, :2].tolist()) == {0, 8} - - -def test_source_node_local_maps_fall_back_to_global_replicas(): - redundant = torch.tensor([[4], [5], [0], [1]], dtype=torch.int64) - maps = [build_logical_to_physical_map(redundant, 8, source_rank=rank, node_world_size=2) for rank in range(4)] - assert maps[0][1][0].item() == maps[1][1][0].item() == 1 - assert maps[2][1][0].item() == maps[3][1][0].item() == 1 - assert maps[0][0][0, 0].item() == maps[1][0][0, 0].item() == 0 - assert maps[2][0][0, 0].item() == maps[3][0][0, 0].item() == 8 - - -def test_source_rank_rotates_selected_replica_order_without_changing_copies(): - # Expert 0 is primary on rank 0 and redundant on rank 1, so both ranks - # on node 0 have the same two local copies. Their source-rank phases - # must differ while their selected set/count remain identical. - redundant = torch.tensor([[1], [0], [3], [2]], dtype=torch.int64) - rank0_map, rank0_count = build_logical_to_physical_map(redundant, 4, source_rank=0, node_world_size=2) - rank1_map, rank1_count = build_logical_to_physical_map(redundant, 4, source_rank=1, node_world_size=2) - - assert rank0_count[0].item() == rank1_count[0].item() == 2 - assert set(rank0_map[0, :2].tolist()) == set(rank1_map[0, :2].tolist()) == {0, 3} - assert torch.equal(rank1_map[0, :2], torch.tensor([3, 0], dtype=torch.int32)) - - -@pytest.mark.parametrize( - "source_rank,node_world_size", - [(None, None), (0, 2), (1, 2), (2, 2), (3, 2)], -) -def test_logical_to_physical_maps_for_layers_match_single_layer_api(source_rank, node_world_size): - placements_by_layer = torch.tensor( - [ - [[4, 5], [0, 1], [0, 1], [2, 3]], - [[6, 7], [0, 1], [0, 1], [2, 3]], - [[4, 5], [0, 1], [0, 1], [2, 3]], - ], - dtype=torch.int64, - ) - - maps_by_layer, counts_by_layer = build_logical_to_physical_maps_for_layers( - placements_by_layer, - num_logical_experts=8, - source_rank=source_rank, - node_world_size=node_world_size, - ) - expected_by_layer = [ - build_logical_to_physical_map( - placement, - num_logical_experts=8, - source_rank=source_rank, - node_world_size=node_world_size, - ) - for placement in placements_by_layer - ] - - assert maps_by_layer.shape == (3, 8, 4) - assert counts_by_layer.shape == (3, 8) - assert maps_by_layer.dtype == counts_by_layer.dtype == torch.int32 - assert torch.equal(maps_by_layer, torch.stack([item[0] for item in expected_by_layer])) - assert torch.equal(counts_by_layer, torch.stack([item[1] for item in expected_by_layer])) - if source_rank is None: - # expert 0 的主副本在 rank 0,且 rank 1、2 都有一个冗余副本;第 1、2 - # 个冗余副本必须分别写入映射表的第 1、2 列,而不能互相覆盖。 - assert torch.equal(counts_by_layer[:, 0], torch.tensor([3, 3, 3], dtype=torch.int32)) - assert torch.equal( - maps_by_layer[:, 0, :3], - torch.tensor([[0, 6, 10], [0, 6, 10], [0, 6, 10]], dtype=torch.int32), - ) - positions = torch.arange(maps_by_layer.shape[-1]).view(1, 1, -1) - valid = positions < counts_by_layer.unsqueeze(-1) - assert torch.all(maps_by_layer[valid] >= 0) - assert torch.all(maps_by_layer[~valid] == -1) - - -def test_plan_redundant_experts_prefers_local_node_load_relief(): - source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) - source_load[0, 0, 0, 0] = 1024 - placement = plan_redundant_experts( - source_load, - num_ranks=4, - num_redundant_experts_per_rank=1, - expert_alignment=128, - node_world_size=2, - ) - assert placement[0, 1, 0] == 0 - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) - assert torch.equal(predicted, _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128)) - - -def test_plan_redundant_experts_single_node_matches_default_behavior(): - load = torch.tensor([[1000, 900, 800, 700, 600, 500, 400, 300]]).unsqueeze(0).unsqueeze(2) - default = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1) - single_node = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=1, node_world_size=4) - assert torch.equal(single_node, default) - - -def test_source_node_estimate_matches_local_first_runtime_replica_sharing(): - # Expert 0 has a copy on each node. Node 0 and node 1 issue unequal - # traffic, so collapsing them before planning produces the wrong result. - placement = torch.tensor([[[4], [5], [0], [1]]], dtype=torch.int64) - source_load = torch.zeros((1, 1, 2, 8), dtype=torch.int64) - source_load[0, 0, 0, 0] = 256 - source_load[0, 0, 1, 0] = 128 - - predicted = _estimate_rank_load(source_load, placement, expert_alignment=128, node_world_size=2) - runtime = _manual_runtime_rank_load(source_load, placement, node_world_size=2, alignment=128) - collapsed_global = _estimate_rank_load(source_load.sum(dim=2, keepdim=True), placement, expert_alignment=128) - - assert torch.equal(predicted, runtime) - assert torch.equal(predicted[0, 0], torch.tensor([256.0, 0.0, 128.0, 0.0])) - assert not torch.equal(predicted, collapsed_global) - - -def test_source_node_planner_constraints_and_real_critical_improvement(): - source_load = torch.tensor( - [ - [ - [ - [697, 451, 383, 536, 349, 404, 854, 425], - [861, 103, 166, 612, 444, 263, 910, 392], - ] - ], - [ - [ - [944, 457, 338, 108, 63, 525, 48, 216], - [439, 117, 837, 550, 833, 201, 729, 5], - ] - ], - [ - [ - [749, 159, 18, 723, 12, 700, 419, 51], - [112, 135, 8, 840, 40, 970, 90, 683], - ] - ], - ], - dtype=torch.int64, - ) - initial = build_initial_redundant_expert_ids(8, 4, 1).unsqueeze(0) - planned = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) - - for rank, experts in enumerate(planned[0].tolist()): - assert len(experts) == len(set(experts)) == 1 - assert experts[0] // 2 != rank - - before = _estimate_rank_load(source_load, initial, expert_alignment=128, node_world_size=2) - after = _estimate_rank_load(source_load, planned, expert_alignment=128, node_world_size=2) - manual_before = _manual_runtime_rank_load(source_load, initial, 2, 128) - manual_after = _manual_runtime_rank_load(source_load, planned, 2, 128) - assert torch.equal(before, manual_before) - assert torch.equal(after, manual_after) - assert after.max(dim=2).values.sum() < before.max(dim=2).values.sum() - - -def test_source_node_select_uses_the_same_runtime_critical_prediction(): - source_load = torch.zeros((2, 1, 2, 8), dtype=torch.int64) - source_load[:, 0, 0, 0] = torch.tensor([1024, 768]) - source_load[:, 0, 1, 6] = torch.tensor([896, 1024]) - current = torch.tensor([[[2], [4], [6], [0]]], dtype=torch.int64) - candidate = plan_redundant_experts(source_load, 4, 1, expert_alignment=128, node_world_size=2) - selected, _improved, _metrics, _before_load, _after_load = select_improving_placements( - source_load, - current, - candidate, - rebalance_gain_threshold=0.05, - expert_alignment=128, - node_world_size=2, - ) - assert torch.equal( - _estimate_rank_load(source_load, selected, 128, 2), - _manual_runtime_rank_load(source_load, selected, 2, 128), - ) - - -def _count_moved_slots(current: torch.Tensor, target: torch.Tensor) -> int: - """Count rank rows gaining an expert: one migrated row per new expert id.""" - moved = 0 - for layer in range(current.shape[0]): - for rank in range(current.shape[1]): - moved += len(set(target[layer, rank].tolist()) - set(current[layer, rank].tolist())) - return moved - - -def test_sticky_plan_reproduces_current_when_load_unchanged(): - generator = torch.Generator().manual_seed(7) - load = torch.randint(1, 1000, (3, 16, 32), generator=generator).unsqueeze(2) - placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - - replanned = plan_redundant_experts( - load, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=placement, - stickiness=0.1, - ) - - assert torch.equal(replanned, placement) - for layer in range(placement.shape[0]): - assert build_transfer_plan(placement[layer], replanned[layer], 32, 4, 4) == [] - - -def test_sticky_plan_bounded_moves_under_small_perturbation(): - generator = torch.Generator().manual_seed(11) - load = torch.randint(100, 1000, (4, 16, 32), generator=generator).unsqueeze(2) - placement = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - noise = torch.rand((4, 16, 32), generator=generator).unsqueeze(2) * 0.1 + 0.95 - perturbed = (load.double() * noise).round().to(torch.int64) - - sticky = plan_redundant_experts(perturbed, 4, 2, current_placement=placement, stickiness=0.1) - free = plan_redundant_experts(perturbed, 4, 2) - - sticky_moves = _count_moved_slots(placement, sticky) - free_moves = _count_moved_slots(placement, free) - assert sticky_moves <= placement.numel() // 4 - assert sticky_moves < free_moves - - def critical(candidate): - return _estimate_rank_load(perturbed, candidate).max(dim=2).values.sum() - - assert critical(sticky) <= critical(free) * 1.1 - - -def test_sticky_plan_still_churns_under_phase_shift(): - layers, experts = 8, 32 - before = torch.full((layers, experts), 10, dtype=torch.int64) - after = torch.full((layers, experts), 10, dtype=torch.int64) - offsets = torch.arange(4) - for layer in range(layers): - before[layer, (4 * layer + offsets) % experts] = 5000 - after[layer, (4 * layer + 16 + offsets) % experts] = 5000 - before = before.unsqueeze(0).unsqueeze(2) - after = after.unsqueeze(0).unsqueeze(2) - placement = plan_redundant_experts(before, num_ranks=4, num_redundant_experts_per_rank=2) - - replanned = plan_redundant_experts( - after, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=placement, - stickiness=0.1, - ) - - assert _count_moved_slots(placement, replanned) > placement.numel() // 2 - - -def test_transfer_plan_slot_permutation_is_free(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 4], [7, 6], [1, 0], [3, 2]]) - - assert torch.equal(align_target_placement(current, target), current) - assert build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) == [] - - -def test_align_target_placement_keeps_retained_experts_in_live_slots(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) - - canonical = align_target_placement(current, target) - - assert torch.equal(canonical, torch.tensor([[6, 5], [0, 7], [0, 1], [2, 3]])) - - -def test_canonical_placement_keeps_transfer_rows_and_published_map_consistent(): - num_logical_experts = 8 - world_size = 4 - num_redundant_slots_per_rank = 2 - num_experts_per_rank = num_logical_experts // world_size - num_physical_experts_per_rank = num_experts_per_rank + num_redundant_slots_per_rank - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = torch.tensor([[5, 6], [7, 0], [1, 0], [3, 2]]) - canonical = align_target_placement(current, target) - plan = build_transfer_plan(current, canonical, num_logical_experts, world_size, node_world_size=2) - - # Label every current physical row by its resident logical expert, then - # apply the transfer plan from a frozen source snapshot just as staging - # copies do before the destination rows are published. - source_rows = [ - list(range(rank * num_experts_per_rank, (rank + 1) * num_experts_per_rank)) + current[rank].tolist() - for rank in range(world_size) - ] - live_rows = [row.copy() for row in source_rows] - for step in plan: - live_rows[step.dst_rank][num_experts_per_rank + step.dst_slot] = source_rows[step.src_rank][step.src_local_row] - - logical_to_physical, replica_count = build_logical_to_physical_map(canonical, num_logical_experts) - for logical_expert, count in enumerate(replica_count.tolist()): - for physical_id in logical_to_physical[logical_expert, :count].tolist(): - rank, row = divmod(physical_id, num_physical_experts_per_rank) - assert live_rows[rank][row] == logical_expert - - -def test_plan_and_broadcast_publishes_canonical_placement(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.world_size = 4 - manager.node_world_size = 2 - manager.num_logical_experts = 8 - manager.num_redundant_experts_per_rank = 2 - manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) - manager.placement_stickiness = 0.1 - manager.rebalance_gain_threshold = 0.05 - manager.evaluation_group = object() - candidate = torch.tensor([[[5, 4], [7, 6], [1, 0], [3, 2]]]) - broadcasts = [] - - def fixed_selector(*_args, **_kwargs): - rank_load = torch.full((1, 1, 4), 100.0) - return candidate.clone(), torch.tensor([True]), {}, rank_load, rank_load - - def record_broadcast(result_list, **_kwargs): - broadcasts.append(result_list[0]) - - monkeypatch.setattr( - manager_module, - "plan_redundant_experts", - lambda *_args, **_kwargs: candidate.clone(), - ) - monkeypatch.setattr(manager_module, "select_improving_placements", fixed_selector) - monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) - - result = manager._plan_and_broadcast(torch.full((1, 1, 2, 8), 100, dtype=torch.int64)) - - assert torch.equal(result["placement"], manager.current_placement) - assert broadcasts and torch.equal(broadcasts[0]["placement"], manager.current_placement) - - -def test_stickiness_zero_matches_unbiased_plan(): - generator = torch.Generator().manual_seed(17) - load = torch.randint(1, 1000, (2, 8, 16), generator=generator).unsqueeze(2) - unbiased = plan_redundant_experts(load, num_ranks=4, num_redundant_experts_per_rank=2) - unrelated = build_initial_redundant_expert_ids(16, 4, 2).unsqueeze(0).expand(8, -1, -1).clone() - - replanned = plan_redundant_experts( - load, - num_ranks=4, - num_redundant_experts_per_rank=2, - current_placement=unrelated, - stickiness=0.0, - ) - - assert torch.equal(replanned, unbiased) - - -def test_plan_and_broadcast_propagates_rank_zero_error_after_existing_broadcast( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.global_rank = 0 - manager.world_size = 2 - manager.node_world_size = 2 - manager.num_logical_experts = 8 - manager.num_redundant_experts_per_rank = 2 - manager.current_placement = torch.tensor([[[4, 5], [6, 7], [0, 1], [2, 3]]]) - manager.placement_stickiness = 0.1 - manager.rebalance_gain_threshold = 0.05 - manager.evaluation_group = object() - broadcasted = [] - - def broken_planner(*_args, **_kwargs): - raise RuntimeError("planner boom") - - def record_broadcast(result_list, **_kwargs): - broadcasted.append(result_list[0]) - - monkeypatch.setattr(manager_module, "plan_redundant_experts", broken_planner) - monkeypatch.setattr(manager_module.dist, "broadcast_object_list", record_broadcast) - - with pytest.raises(RuntimeError, match="EPLB planner failed on rank zero") as exc_info: - manager._plan_and_broadcast(torch.full((1, 1, 1, 8), 100, dtype=torch.int64)) - - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "planner boom" - assert broadcasted == [{"kind": "error", "message": "RuntimeError: planner boom"}] - - -def test_steady_state_sparse_sampling_records_steps_sixteen_to_nineteen_and_evaluates_step_twenty( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=1, - num_logical_experts=4, - world_size=1, - ), - }, - )() - ] - manager.in_flight = False - manager.prefill_steps = 15 - manager.step_interval = 20 - manager.sampling_interval = manager.step_interval - manager.evaluation_group = object() - manager.num_logical_experts = 4 - manager.global_rank = 1 - manager.evaluation_in_flight = False - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - - recordings, resets, started = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - monkeypatch.setattr( - manager_module.torch.cuda, - "synchronize", - lambda: pytest.fail("step must not synchronize CUDA"), - ) - - manager.step() - assert manager.prefill_steps == 16 - assert recordings == [True] - assert resets == [True] - assert manager._steady_collection_end_step == 20 - - for _ in range(3): - manager.step() - assert manager.prefill_steps == 19 - assert recordings == [True] - assert started == [] - - manager.step() - assert manager.prefill_steps == 20 - assert manager._steady_collection_end_step is None - assert started == [True] - - -def test_steady_sampling_window_clamps_to_short_interval_without_moving_boundary( - monkeypatch, -): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = False - manager.prefill_steps = 0 - manager.step_interval = 20 - manager.sampling_interval = 3 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - recordings, resets, started = [], [], [] - manager._set_recording = lambda enabled: recordings.append((manager.prefill_steps, enabled)) - manager._reset_recorded_samples = lambda: resets.append(manager.prefill_steps) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(manager.prefill_steps)) - - manager._prepare_next_sampling_window() - assert resets == [0] - assert recordings == [(0, True)] - assert manager._steady_collection_end_step == 3 - - manager.step() - manager.step() - assert started == [] - manager.step() - assert started == [3] - - -def test_eplb_step_does_not_start_a_second_evaluation_while_one_is_pending(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - counter = torch.tensor([[10, 20, 30, 40], [0, 0, 0, 0]], dtype=torch.int64) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=1, - num_logical_experts=4, - world_size=1, - ), - }, - )() - ] - manager.in_flight = False - manager.prefill_steps = 1 - manager.step_interval = 2 - manager.sampling_interval = manager.step_interval - manager.evaluation_group = object() - manager.num_logical_experts = 4 - manager.global_rank = 1 - manager.evaluation_in_flight = True - - started = [] - monkeypatch.setattr(manager, "_poll_evaluation", lambda: True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - - manager.step() - - assert manager.prefill_steps == 1 - assert started == [] - - -def test_evaluation_no_improvement_logs_model_fields_without_reopening_interval_window( - monkeypatch, -): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 0 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.weights = [] - manager._eplb_states = [] - recordings, logs = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) - - assert not manager._poll_evaluation() - assert recordings == [False] - assert "model_imbalance_ratio" in logs[0][0] - assert "candidate_rebalance_gain" in logs[0][0] - assert "candidate_changed_layer_count" in logs[0][0] - assert "actual_changed_layer_count" in logs[0][0] - assert "next_sampling_interval" in logs[0][0] - assert manager.sampling_interval == 80 - - -def test_interval_one_rearms_after_evaluation_but_never_evaluates_empty_counter( - monkeypatch, -): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.prefill_steps = 1 - manager.step_interval = 1 - manager.sampling_interval = 1 - manager._steady_collection_end_step = None - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 1 - manager.weights = [] - manager._eplb_states = [] - recordings, starts = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._start_evaluation = lambda: starts.append(True) - manager._evaluation_ready_on_all_ranks = lambda: True - - manager.poll() - - # The no-improvement backoff changes interval 1 to 4. The clamped - # steady window arms immediately but still waits for boundary step 5. - assert recordings == [True] - assert manager._steady_collection_end_step is not None - assert manager._steady_collection_end_step == 5 - assert starts == [] - assert manager.prefill_steps == 1 - - for _ in range(3): - manager.step() - assert starts == [] - manager.step() - assert starts == [True] - - -def test_evaluation_worker_error_is_raised_by_main_thread(): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = RuntimeError("planner failed") - manager._evaluation_result = None - manager._evaluation_thread = DoneThread() - - with pytest.raises(RuntimeError, match="planner failed"): - manager._poll_evaluation() - - -def test_evaluation_state_is_cleared_before_second_round(monkeypatch): - class DoneThread: - def join(self): - pass - - class PendingThread: - def __init__(self, **_kwargs): - self.started = False - - def start(self): - self.started = True - - def join(self): - pytest.fail("pending worker must not be joined") - - class Event: - def record(self, _stream): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 1, - } - manager.global_rank = 1 - manager.prefill_steps = 0 - manager.step_interval = 1 - manager.sampling_interval = 1 - manager.weights = [] - manager._eplb_states = [] - manager._set_recording = lambda _enabled: None - monkeypatch.setattr(manager_module.threading, "Thread", PendingThread) - monkeypatch.setattr(manager_module.torch.cuda, "Event", Event) - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", lambda: object()) - - assert not manager._poll_evaluation() - assert manager._evaluation_result is None - assert manager._evaluation_error is None - - manager._start_evaluation() - assert manager.evaluation_in_flight - assert manager._evaluation_result is None - assert manager._poll_evaluation() # New worker has not produced a result. - - -def test_manager_collects_recent_ring_samples_in_chronological_order(monkeypatch): - counters = [ - torch.tensor([[10, 11], [20, 21], [30, 31]], dtype=torch.int64), - torch.tensor([[40, 41], [50, 51], [60, 61]], dtype=torch.int64), - ] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=5, - num_logical_experts=2, - world_size=1, - ), - }, - )() - for counter in counters - ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] - manager.evaluation_group = object() - metadata_sizes = [] - - def all_reduce(metadata, **_kwargs): - metadata_sizes.append(metadata.numel()) - - monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) - - samples = manager._collect_local_samples() - - assert metadata_sizes == [4] - assert torch.equal( - samples, - torch.tensor( - [ - [[30, 31], [60, 61]], - [[10, 11], [40, 41]], - [[20, 21], [50, 51]], - ], - dtype=torch.int64, - ), - ) - - -def test_manager_collects_only_two_recent_sparse_samples(monkeypatch): - counters = [ - torch.tensor([[10, 11], [20, 21], [30, 31], [40, 41]], dtype=torch.int64), - torch.tensor([[50, 51], [60, 61], [70, 71], [80, 81]], dtype=torch.int64), - ] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=2, - num_logical_experts=2, - world_size=1, - ), - }, - )() - for counter in counters - ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] - manager.evaluation_group = object() - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **kwargs: None) - - samples = manager._collect_local_samples() - - assert torch.equal( - samples, - torch.tensor([[[10, 11], [50, 51]], [[20, 21], [60, 61]]], dtype=torch.int64), - ) - - -def test_eplb_counter_capacity_covers_default_dense_interval(monkeypatch): - args = type( - "Args", - (), - {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, - )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - monkeypatch.setattr(fused_weight_module, "get_prefill_eplb_step_interval", lambda: 20) - monkeypatch.setattr(fused_weight_module, "get_node_world_size", lambda: 2) - monkeypatch.setattr(torch.Tensor, "cuda", lambda tensor: tensor) - original_zeros = torch.zeros - - def cpu_zeros(*shape, **kwargs): - kwargs.pop("device", None) - return original_zeros(*shape, **kwargs) - - monkeypatch.setattr(fused_weight_module.torch, "zeros", cpu_zeros) - - weight._init_expert_parallel_state() - - assert weight.expert_parallel_state.eplb.route_counter.shape == (40, 4) - - -def test_steady_sampling_resets_fixed_ring_without_retained_history(): - counter = torch.ones((8, 4), dtype=torch.int64) - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - state = _test_parallel_state( - eplb=True, - route_counter=counter, - recording=False, - recorded_sample_count=99, - num_logical_experts=4, - world_size=1, - ).eplb - manager._eplb_states = [state] - - manager._reset_recorded_samples() - manager._reset_recorded_samples() - - assert state.route_counter.shape == (8, 4) - assert torch.count_nonzero(state.route_counter) == 0 - assert state.recorded_sample_count == 0 - assert not hasattr(manager, "_retained_local_samples") - assert not hasattr(manager, "_sample_history") - - -def test_ep_without_eplb_creates_layout_without_eplb_runtime_state(monkeypatch): - args = type( - "Args", - (), - {"enable_prefill_eplb": False, "eplb_num_redundant_experts_per_rank": 2}, - )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - - weight._init_expert_parallel_state() - - assert weight.expert_parallel_state is not None - assert weight.expert_parallel_state.eplb is None - assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts - assert weight._initial_redundant_expert_ids == [] - assert not hasattr(weight, "route_counter") - assert not hasattr(weight, "routed_expert_counter_tensor") - - -def test_disable_eplb_model_init_skips_eplb_state(monkeypatch): - args = type( - "Args", - (), - {"enable_prefill_eplb": True, "eplb_num_redundant_experts_per_rank": 2}, - )() - weight = object.__new__(fused_weight_module.FusedMoeWeight) - weight.n_routed_experts = 4 - weight.global_world_size = 2 - weight.global_rank_ = 0 - weight.enable_ep_moe = True - monkeypatch.setattr(fused_weight_module, "get_env_start_args", lambda: args) - monkeypatch.setattr( - fused_weight_module, - "build_initial_redundant_expert_ids", - lambda *args, **kwargs: pytest.fail("disabled scope must not initialize EPLB"), - ) - - with disable_eplb_model_init(): - weight._init_expert_parallel_state() - - assert weight.expert_parallel_state.eplb is None - assert weight.expert_parallel_state.num_total_physical_experts == weight.n_routed_experts - assert weight._initial_redundant_expert_ids == [] - - -def test_disable_eplb_model_init_scope_restores_after_exception(): - assert not is_eplb_model_init_disabled() - with disable_eplb_model_init(): - assert is_eplb_model_init_disabled() - with disable_eplb_model_init(): - assert is_eplb_model_init_disabled() - assert is_eplb_model_init_disabled() - - with pytest.raises(RuntimeError): - with disable_eplb_model_init(): - raise RuntimeError - assert not is_eplb_model_init_disabled() - - -def test_manager_evaluation_collective_preserves_source_node_axis(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=torch.zeros((1, 4), dtype=torch.int64), - num_logical_experts=4, - world_size=1, - ) - }, - )() - ] - manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] - manager.global_rank = 2 - manager.world_size = 4 - manager.node_world_size = 2 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.num_logical_experts = 4 - manager.num_redundant_experts_per_rank = 1 - manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0) - manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None - manager._continuous_collection_end_step = None - local = torch.full((1, 1, 4), 100, dtype=torch.int64) - manager._collect_local_samples = lambda: local - seen = {} - - def all_reduce(tensor, **kwargs): - seen["before"] = tensor.clone() - seen["group"] = kwargs["group"] - # Simulate source node 0's contribution from the other ranks. - tensor[:, :, 0].fill_(100) - - monkeypatch.setattr(manager_module.dist, "all_reduce", all_reduce) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) - - def plan_and_broadcast(global_load): - seen["global_load"] = global_load.clone() - return {"kind": "insufficient"} - - manager._plan_and_broadcast = plan_and_broadcast - - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) - - assert seen["group"] is manager.evaluation_group - expected_local = local - assert seen["before"].shape == (1, 1, 2, 4) - assert torch.equal(seen["before"][:, :, 0], torch.zeros_like(expected_local)) - assert torch.equal(seen["before"][:, :, 1], expected_local) - assert torch.equal(seen["global_load"][:, :, 0], torch.full_like(expected_local, 100)) - assert torch.equal(seen["global_load"][:, :, 1], expected_local) - assert manager._evaluation_error is None - assert manager._evaluation_result["recorded_sample_count"] == 1 - assert manager._evaluation_result["sample_window_steps"] == 4 - - -def test_manager_planned_evaluation_builds_improved_metadata_in_one_multilayer_call(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=torch.zeros((1, 4), dtype=torch.int64), - num_logical_experts=4, - world_size=4, - ) - }, - )() - for _ in range(3) - ] - manager._eplb_states = [weight.expert_parallel_state.eplb for weight in manager.weights] - manager.global_rank = 1 - manager.world_size = 4 - manager.node_world_size = 2 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.num_logical_experts = 4 - manager.num_redundant_experts_per_rank = 1 - manager.current_placement = build_initial_redundant_expert_ids(4, 4, 1).unsqueeze(0).expand(3, -1, -1).clone() - manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None - manager._continuous_collection_end_step = None - prepare_calls = [] - prepared_batches = object() - - def prepare_transfer(layer_plans): - prepare_calls.append(layer_plans) - return prepared_batches - - manager.transfer = SimpleNamespace(prepare_transfer=prepare_transfer) - manager._collect_local_samples = lambda: torch.full((1, 3, 4), 100, dtype=torch.int64) - planned_placement = torch.tensor( - [ - [[3], [0], [1], [2]], - [[2], [3], [0], [1]], - [[1], [2], [3], [0]], - ], - dtype=torch.int64, - ) - manager._plan_and_broadcast = lambda _global_load: { - "kind": "planned", - "placement": planned_placement, - "improved": torch.tensor([True, False, True]), - } - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) - calls = [] - original_build_maps_for_layers = manager_module.build_logical_to_physical_maps_for_layers - - def build_maps_for_layers(*args, **kwargs): - calls.append(args[0].shape) - return original_build_maps_for_layers(*args, **kwargs) - - monkeypatch.setattr( - manager_module, - "build_logical_to_physical_maps_for_layers", - build_maps_for_layers, - ) - - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) - - assert manager._evaluation_error is None - assert calls == [torch.Size([2, 4, 1])] - metadata = manager._evaluation_result["metadata"] - assert metadata[1] is None - assert [layer_index for layer_index, _plan in manager._evaluation_result["layer_plans"]] == [0, 2] - assert prepare_calls == [manager._evaluation_result["layer_plans"]] - assert manager._evaluation_result["prepared_batches"] is prepared_batches - for layer_index in (0, 2): - item = metadata[layer_index] - expected = build_logical_to_physical_map( - planned_placement[layer_index], - 4, - source_rank=manager.global_rank, - node_world_size=manager.node_world_size, - ) - assert torch.equal(item[0], expected[0]) - assert torch.equal(item[1], expected[1]) - - -def test_manager_preparation_error_is_saved_as_evaluation_error(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.weights = [ - type( - "Weight", - (), - { - "expert_parallel_state": _test_parallel_state( - eplb=True, - route_counter=torch.zeros((1, 2), dtype=torch.int64), - num_logical_experts=2, - world_size=1, - ) - }, - )() - ] - manager._eplb_states = [manager.weights[0].expert_parallel_state.eplb] - manager.global_rank = 0 - manager.world_size = 1 - manager.node_world_size = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.num_logical_experts = 2 - manager.current_placement = torch.tensor([[[0], [1]]], dtype=torch.int64) - manager.evaluation_group = object() - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None - manager._continuous_collection_end_step = None - manager._collect_local_samples = lambda: torch.ones((1, 1, 2), dtype=torch.int64) - manager._plan_and_broadcast = lambda _global_load: { - "kind": "planned", - "placement": torch.tensor([[[0], [1]]], dtype=torch.int64), - "improved": torch.tensor([True]), - } - - def fail_prepare(_layer_plans): - raise RuntimeError("prepare failed") - - manager.transfer = SimpleNamespace(prepare_transfer=fail_prepare) - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda _tensor, **_kwargs: None) - monkeypatch.setattr(manager_module.torch.cuda, "set_device", lambda _device: None) - monkeypatch.setattr(manager_module, "build_transfer_plan", lambda *_args, **_kwargs: []) - - manager._evaluate_after_event(type("Event", (), {"synchronize": lambda self: None})()) - - assert manager._evaluation_result is None - assert isinstance(manager._evaluation_error, RuntimeError) - assert str(manager._evaluation_error) == "prepare failed" - - -def test_decode_dispatch_keeps_logical_ids_and_uses_logical_expert_count(monkeypatch): - class Buffer: - def low_latency_dispatch(self, **kwargs): - calls.append(kwargs) - return "recv", "masked", "handle", "event", "hook" - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.quant_method = type("Quant", (), {"method_name": "fp8"})() - impl.n_routed_experts = 128 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - impl._select_experts = lambda **_kwargs: ( - torch.ones((1, 2)), - torch.tensor([[0, 127]], dtype=torch.int32), - torch.tensor([[0, 127]], dtype=torch.int32), - ) - calls = [] - monkeypatch.setattr( - deepgemm_module, - "get_deepep_num_max_dispatch_tokens_per_rank_decode", - lambda: 16, - ) - monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_low_latency_buffer", Buffer()) - - result = impl.low_latency_dispatch( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - False, - 2, - False, - 0, - 0, - "softmax", - ) - - assert result[2].tolist() == [[0, 127]] - assert calls[0]["num_experts"] == 128 - - -def test_decode_select_does_not_clone_or_map_eplb_topk_ids(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import topk_select - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - topk_ids = torch.tensor([[3, 127]], dtype=torch.int32) - monkeypatch.setattr(topk_select, "select_experts", lambda **_kwargs: (torch.ones((1, 2)), topk_ids)) - _, selected, origin = impl._select_experts( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - 2, - False, - False, - 0, - 0, - "softmax", - is_prefill=False, - ) - - assert selected.data_ptr() == origin.data_ptr() - - -def test_eplb_prefill_uses_single_fused_path_for_global_topk(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 - impl.quant_method = object() - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - physical_ids = torch.tensor([[130, 131]], dtype=torch.long) - calls = [] - - def fused_topk(**kwargs): - calls.append(kwargs) - return torch.ones((1, 2)), physical_ids, None - - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) - monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") - - weights, topk_idx, qinput = impl.select_experts_and_quant_input( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - object(), - False, - 2, - False, - 0, - 0, - "softmax", - ) - - assert weights.tolist() == [[1.0, 1.0]] - assert topk_idx is physical_ids - assert topk_idx.dtype is torch.long - assert qinput == "qinput" - assert not calls[0]["use_grouped_topk"] - assert not calls[0]["return_logical_ids"] - - -def test_eplb_prefill_dispatch_consumes_physical_ids_and_event(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk - - class Buffer: - def dispatch(self, _qinput, **kwargs): - calls.append(kwargs) - return ( - (torch.empty((4, 2)),), - "recv_idx", - "recv_weight", - SimpleNamespace(num_recv_tokens_per_expert_list=[4]), - SimpleNamespace(current_stream_wait=lambda: None), - ) - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 - impl.quant_method = object() - state = _test_parallel_state( - eplb=True, - route_counter=torch.zeros((3, 128), dtype=torch.int64), - recording=True, - ) - _set_expert_parallel_state(impl, state) - impl.ep_balance_counters = None - calls, fused_calls = [], [] - physical_ids = torch.tensor([[130, 131]], dtype=torch.long) - - def fused_topk(**kwargs): - fused_calls.append(kwargs) - return torch.ones((1, 2)), physical_ids, None - - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) - monkeypatch.setattr(deepgemm_module, "quantize_fused_experts_input", lambda *_args: "qinput") - monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) - monkeypatch.setattr( - deepgemm_module, - "get_deepep_num_max_dispatch_tokens_per_rank_prefill", - lambda: 16, - ) - monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) - - weights, topk_idx, qinput = impl.select_experts_and_quant_input( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - object(), - True, - 2, - False, - 1, - 8, - "sigmoid", - ) - caller_event = object() - impl.dispatch( - qinput, - topk_idx, - weights, - overlap_event=caller_event, - ) - - assert topk_idx is physical_ids - assert len(fused_calls) == 1 - assert fused_calls[0]["sample_index"] == 0 - assert fused_calls[0]["record_load"] - assert state.eplb.recorded_sample_count == 1 - assert calls[0]["topk_idx"] is physical_ids - assert calls[0]["topk_idx"].dtype is torch.long - assert calls[0]["previous_event"] is caller_event - - -def test_prefill_dispatch_preserves_event(monkeypatch): - class Buffer: - def dispatch(self, _qinput, **kwargs): - calls.append(kwargs) - return ( - (torch.empty((4, 2)),), - "recv_idx", - "recv_weight", - SimpleNamespace(num_recv_tokens_per_expert_list=[4]), - SimpleNamespace(current_stream_wait=lambda: None), - ) - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - impl.ep_balance_counters = None - calls = [] - caller_event = object() - monkeypatch.setattr(deepgemm_module.dist_group_manager, "ep_buffer", Buffer()) - monkeypatch.setattr( - deepgemm_module, - "get_deepep_num_max_dispatch_tokens_per_rank_prefill", - lambda: 16, - ) - monkeypatch.setattr(deepgemm_module, "get_ep_num_sms", lambda: 8) - - impl.dispatch( - "qinput", - torch.tensor([[1, 2]], dtype=torch.long), - torch.ones((1, 2)), - caller_event, - ) - - assert calls[0]["previous_event"] is caller_event - assert calls[0]["topk_idx"].dtype is torch.long - - -def test_deepgemm_constructor_configures_eplb(): - state = _validated_expert_parallel_state() - impl = deepgemm_module.FuseMoeDeepGEMM(4, 0, 1.0, SimpleNamespace(), expert_parallel_state=state) - assert impl.expert_parallel_state is state - - -def test_prefill_eplb_returns_requested_logical_ids(monkeypatch): - from lightllm.common.basemodel.triton_kernel.fused_moe import grouped_topk - - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.routed_scaling_factor = 1.0 - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True)) - - def fused_topk(**kwargs): - assert kwargs["return_logical_ids"] - return ( - torch.ones((1, 2)), - torch.tensor([[13, 14]], dtype=torch.int32), - torch.tensor([[3, 4]], dtype=torch.int32), - ) - - monkeypatch.setattr(grouped_topk, "triton_grouped_topk_eplb", fused_topk) - _, physical_ids, logical_ids = impl._select_experts( - torch.empty((1, 4)), - torch.empty((1, 128)), - None, - 2, - False, - False, - 0, - 0, - "softmax", - is_prefill=True, - preserve_logical_ids=True, - ) - - assert physical_ids.tolist() == [[13, 14]] - assert logical_ids.tolist() == [[3, 4]] - - -def test_decode_masked_group_gemm_uses_primary_rows_only_when_eplb_is_enabled( - monkeypatch, -): - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - _set_expert_parallel_state(impl, _test_parallel_state(eplb=True, num_logical_experts=8, world_size=1)) - captured = {} - - def masked(*args, **kwargs): - captured["w13"] = args[3] - captured["w13_scale"] = args[4] - captured["w2"] = args[5] - captured["w2_scale"] = args[6] - return "out" - - monkeypatch.setattr(deepgemm_module, "masked_group_gemm", masked) - pack = lambda: type( - "Pack", - (), - {"weight": torch.empty((10, 4)), "weight_scale": torch.empty((10, 1))}, - )() - - assert impl.masked_group_gemm((torch.empty((1, 4)),), pack(), pack(), torch.empty(8), torch.float16, 1) == "out" - assert captured["w13"].shape[0] == captured["w2"].shape[0] == 8 - assert captured["w13_scale"].shape[0] == captured["w2_scale"].shape[0] == 8 - - -def test_decode_fused_experts_uses_cached_primary_weight_packs_and_logical_experts( - monkeypatch, -): - impl = object.__new__(deepgemm_module.FuseMoeDeepGEMM) - impl.n_routed_experts = 128 - _set_expert_parallel_state( - impl, - _test_parallel_state(eplb=True, num_logical_experts=128, world_size=16, num_redundant_experts_per_rank=2), - ) - impl.quant_method = object() - impl.ep_balance_counters = None - captured = [] - - def fused(**kwargs): - captured.append(kwargs) - return "out" - - monkeypatch.setattr(deepgemm_module, "fused_experts", fused) - pack = lambda: type( - "Pack", - (), - { - "weight": torch.empty((10, 4)), - "weight_scale": torch.empty((10, 1)), - "weight_zero_point": None, - }, - )() - w13, w2 = pack(), pack() - - for _ in range(2): - assert ( - impl._fused_experts( - torch.empty((1, 4)), - w13, - w2, - torch.ones((1, 2)), - torch.zeros((1, 2), dtype=torch.int64), - is_prefill=False, - ) - == "out" - ) - - assert [call["num_experts"] for call in captured] == [128, 128] - assert all(call["w13"].weight.shape[0] == call["w2"].weight.shape[0] == 8 for call in captured) - assert captured[0]["w13"] is captured[1]["w13"] - assert captured[0]["w2"] is captured[1]["w2"] - - -def test_transfer_plan_uses_existing_rows_and_prefers_local_node_replicas(): - current = torch.tensor([[4, 5], [6, 7], [0, 1], [2, 3]]) - target = current.clone() - target[0, 0] = 6 # primary r3, but r1 replica is on r0's node. - target[2, 1] = 4 # primary r2 is local to destination r2. - plan = build_transfer_plan(current, target, num_logical_experts=8, world_size=4, node_world_size=2) - by_dst = {(step.dst_rank, step.dst_slot): step for step in plan} - assert len(by_dst) == 2 - assert by_dst[0, 0] == TransferStep(0, 0, 1, 2) - assert by_dst[2, 1] == TransferStep(2, 1, 2, 0) - - -def test_transfer_plan_cross_node_and_stable_source_load_tie_break(): - current = torch.tensor([[0, 1], [2, 3], [4, 5], [4, 7]]) - target = current.clone() - target[0, 0] = 4 - target[0, 1] = 4 - first = build_transfer_plan(current, target, 8, 4, 2) - second = build_transfer_plan(current, target, 8, 4, 2) - assert first == second - selected = [step for step in first if step.dst_rank == 0] - assert [(step.src_rank, step.src_local_row) for step in selected] == [ - (2, 0), - (3, 2), - ] - - -def test_extract_expert_tensors_includes_weight_and_scale_in_order(): - class Pack: - def __init__(self, offset, scale=True): - self.weight = torch.full((3, 2), offset) - self.weight_scale = torch.full((3, 1), offset + 1) if scale else None - - weight = type("Weight", (), {"w13": Pack(1), "w2": Pack(10, scale=False)})() - tensors = extract_eplb_expert_tensors(weight) - assert [name for name, _ in tensors] == [ - "w13.weight", - "w13.weight_scale", - "w2.weight", - ] - - -def test_commit_staging_rows_only_overwrites_redundant_rows(): - live = torch.arange(20).reshape(5, 4) - staging = torch.full((2, 4), -1) - _commit_staging_rows( - live, - staging, - num_experts_per_rank=3, - changed_dst_slots=(0, 1), - ) - assert torch.equal(live[:3], torch.arange(12).reshape(3, 4)) - assert torch.equal(live[3:], staging) - - -def test_commit_staging_rows_preserves_unchanged_destination_slots(): - live = torch.arange(28).reshape(7, 4) - staging = torch.tensor([[-1, -1, -1, -1], [-2, -2, -2, -2], [-3, -3, -3, -3], [-4, -4, -4, -4]]) - original = live.clone() - - _commit_staging_rows( - live, - staging, - num_experts_per_rank=3, - changed_dst_slots=(3, 1), - ) - - assert torch.equal(live[:4], original[:4]) - assert torch.equal(live[4], staging[1]) - assert torch.equal(live[5], original[5]) - assert torch.equal(live[6], staging[3]) - - -def test_commit_staging_rows_merges_contiguous_changed_slots(): - copies = [] - - class View: - def __init__(self, owner, start, length): - self.owner = owner - self.start = start - self.length = length - - def copy_(self, source, **_kwargs): - copies.append( - ( - self.owner, - self.start, - self.length, - source.owner, - source.start, - source.length, - ) - ) - - class Tensor: - def __init__(self, owner, rows): - self.owner = owner - self.shape = (rows,) - - def narrow(self, _dim, start, length): - return View(self.owner, start, length) - - _commit_staging_rows( - Tensor("live", 20), - Tensor("staging", 4), - num_experts_per_rank=10, - changed_dst_slots=(3, 1, 2), - ) - - assert copies == [("live", 11, 3, "staging", 1, 3)] - - -def test_manager_inflight_ready_gate_commits_ordered_prefix_and_propagates_worker_error( - monkeypatch, -): - class Transfer: - def __init__(self): - self.pending = [(0, 0), (1, 1), (2, 2)] - self.commits = [] - self.finished = 0 - - def pending_layers(self): - return self.pending - - def commit(self, layer, buffer_index, post_copy=None): - assert self.pending.pop(0) == (layer, buffer_index) - self.commits.append((layer, buffer_index)) - if post_copy is not None: - post_copy() - - def finish(self): - self.finished += 1 - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer() - manager.in_flight = True - manager.world_size = 2 - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0, 1, 2] - manager._commit_layer_metadata = lambda layer: committed.append(layer) - manager._finish_rebalance = lambda: finished.append(True) - committed, finished = [], [] - operations = [] - current_stream_calls = [] - - class CurrentStream: - def wait_stream(self, stream): - operations.append(("wait", stream)) - - overlap_stream = object() - - def current_stream(): - current_stream_calls.append(True) - return CurrentStream() - - monkeypatch.setattr(manager_module.torch.cuda, "current_stream", current_stream) - monkeypatch.setattr(g_infer_context, "get_overlap_stream", lambda: overlap_stream) - - original_commit = manager.transfer.commit - - def record_commit(*args, **kwargs): - operations.append(("commit", args[0])) - return original_commit(*args, **kwargs) - - manager.transfer.commit = record_commit - - def set_global_ready(count): - return lambda tensor, **kwargs: tensor.fill_(count) - - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(0)) - manager._poll_in_flight() - assert manager.transfer.commits == [] - assert operations == [] - assert current_stream_calls == [] - - # Local rank has three prefetched layers, but global MIN-ready only permits two. - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(2)) - manager._poll_in_flight() - assert manager.transfer.commits == [(0, 0), (1, 1)] - assert committed == [0, 1] - assert not finished - assert manager.transfer.finished == 0 - assert operations == [("wait", overlap_stream), ("commit", 0), ("commit", 1)] - assert current_stream_calls == [True] - - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) - manager._poll_in_flight() - assert finished == [True] - assert manager.transfer.finished == 1 - assert operations == [ - ("wait", overlap_stream), - ("commit", 0), - ("commit", 1), - ("wait", overlap_stream), - ("commit", 2), - ] - assert current_stream_calls == [True, True] - - manager.in_flight_layers = [3] - manager.transfer.pending = [(9, 0)] - monkeypatch.setattr(manager_module.dist, "all_reduce", set_global_ready(1)) - with pytest.raises(RuntimeError, match="does not match expected"): - manager._poll_in_flight() - - class BrokenTransfer: - def pending_layers(self): - raise RuntimeError("boom") - - manager.transfer = BrokenTransfer() - encoded_statuses = [] - - def retain_local_error(tensor, **_kwargs): - encoded_statuses.append(int(tensor.item())) - - monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) - with pytest.raises(RuntimeError, match="EPLB transfer worker failed on this rank") as exc_info: - manager._poll_in_flight() - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "boom" - assert encoded_statuses == [manager_module.EPLB_CONTROL_ERROR] - - -def test_manager_inflight_remote_worker_error_does_not_commit(monkeypatch): - class Transfer: - def __init__(self): - self.commits = [] - - def pending_layers(self): - return [(0, 0)] - - def commit(self, *args): - self.commits.append(args) - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer() - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0] - manager._commit_layer_metadata = lambda _layer: None - manager._finish_rebalance = lambda: None - statuses = [] - - def remote_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) - tensor.fill_(manager_module.EPLB_CONTROL_ERROR) - - monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) - - with pytest.raises(RuntimeError, match="EPLB transfer worker failed on another rank"): - manager._poll_in_flight() - assert statuses == [1] - assert manager.transfer.commits == [] - - -def test_evaluation_ready_gate_propagates_local_and_remote_errors(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager._evaluation_lock = threading.Lock() - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager._evaluation_error = RuntimeError("evaluation boom") - manager._evaluation_result = None - statuses = [] - - def retain_local_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) - - monkeypatch.setattr(manager_module.dist, "all_reduce", retain_local_error) - with pytest.raises(RuntimeError, match="EPLB evaluation failed on this rank") as exc_info: - manager._evaluation_ready_on_all_ranks() - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "evaluation boom" - assert statuses == [manager_module.EPLB_CONTROL_ERROR] - - manager._evaluation_error = None - manager._evaluation_result = {"kind": "no_improvement"} - statuses.clear() - - def remote_error(tensor, **_kwargs): - statuses.append(int(tensor.item())) - tensor.fill_(manager_module.EPLB_CONTROL_ERROR) - - monkeypatch.setattr(manager_module.dist, "all_reduce", remote_error) - with pytest.raises(RuntimeError, match="EPLB evaluation failed on another rank"): - manager._evaluation_ready_on_all_ranks() - assert statuses == [1] - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_manager_inflight_commit_orders_live_weights_between_overlap_forwards( - monkeypatch, -): - class Transfer: - def __init__(self, live, staging): - self.live = live - self.staging = staging - self.pending = [(0, 0)] - - def pending_layers(self): - return self.pending - - def commit(self, layer, buffer_index, post_copy=None): - assert self.pending.pop(0) == (layer, buffer_index) - self.live.copy_(self.staging, non_blocking=True) - if post_copy is not None: - post_copy() - - def finish(self): - pass - - live = torch.tensor([1.0], device="cuda") - staging = torch.tensor([2.0], device="cuda") - previous_read = torch.empty_like(live) - next_read = torch.empty_like(live) - source_stream = torch.cuda.Stream(device=live.device) - destination_stream = torch.cuda.Stream(device=live.device) - initial_stream = torch.cuda.current_stream(device=live.device) - original_overlap_stream = g_infer_context.overlap_stream - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.transfer = Transfer(live, staging) - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - manager.in_flight_layers = [0] - manager._commit_layer_metadata = lambda _layer: None - manager._finish_rebalance = lambda: None - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) - - try: - g_infer_context.overlap_stream = source_stream - with torch.cuda.stream(source_stream): - source_stream.wait_stream(initial_stream) - torch.cuda._sleep(20_000_000) - previous_read.copy_(live, non_blocking=True) - with torch.cuda.stream(destination_stream): - manager._poll_in_flight() - with torch.cuda.stream(source_stream): - source_stream.wait_stream(destination_stream) - next_read.copy_(live, non_blocking=True) - source_stream.synchronize() - - assert previous_read.item() == 1.0 - assert next_read.item() == 2.0 - finally: - g_infer_context.overlap_stream = original_overlap_stream - - -def test_transfer_ring_reuses_a_buffer_only_after_commit_and_consumption(monkeypatch): - operations = [] - - class Event: - def __init__(self): - self.recorded = 0 - self.synchronized = 0 - - def record(self, stream): - self.recorded += 1 - - def synchronize(self): - self.synchronized += 1 - operations.append("consumed synchronize") - - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer.backend = "test" - transfer.device = torch.device("cuda", 0) - transfer.staging_depth = 2 - transfer.staging = [[], []] - transfer.live = [[], [], []] - transfer.num_experts_per_rank = 0 - transfer._release = [threading.Event(), threading.Event()] - for release in transfer._release: - release.set() - transfer._consumed_events = [Event(), Event()] - transfer._consumed_recorded = [False, False] - transfer._changed_dst_slots = [(), ()] - transfer._pending = deque() - transfer._pending_lock = threading.Lock() - transfer._error = None - transfer._thread = None - transfer._needs_staging_reuse_barrier = True - transfer._start_transfer_generation = lambda: None - transfer._finish_transfer_generation = lambda: None - transfer.transfer_group = "transfer-group" - copied = [] - - def copy_batch(batch, _prepared_batch): - for layer, _plan, _buffer, _staging in batch: - copied.append(layer) - operations.append(("copy", layer)) - - transfer._copy_batch = copy_batch - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda device: None) - monkeypatch.setattr(transfer_module.torch.cuda, "Event", Event) - monkeypatch.setattr(transfer_module.torch.cuda, "current_stream", lambda: object()) - monkeypatch.setattr( - transfer_module.dist, - "barrier", - lambda **kwargs: operations.append(("barrier", kwargs["group"])), - ) - - plans = [(0, []), (1, []), (2, [])] - prepared_batches = [(batch, None) for batch in transfer._make_batches(plans)] - monkeypatch.setattr(transfer, "_make_batches", lambda _plans: pytest.fail("start must reuse prepared batches")) - transfer.start(plans, prepared_batches) - deadline = time.monotonic() + 2 - while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: - time.sleep(0.001) - assert transfer.pending_layers() == [(0, 0), (1, 1)] - assert copied == [0, 1] - assert operations == [("copy", 0), ("copy", 1)] - - transfer.commit(0, 0) - while len(transfer.pending_layers()) < 2 and time.monotonic() < deadline: - time.sleep(0.001) - assert transfer.pending_layers() == [(1, 1), (2, 0)] - assert copied == [0, 1, 2] - assert transfer._consumed_events[0].synchronized == 1 - assert operations == [ - ("copy", 0), - ("copy", 1), - "consumed synchronize", - ("barrier", "transfer-group"), - ("copy", 2), - ] - transfer.commit(1, 1) - transfer.commit(2, 0) - transfer.finish() - - -def test_transfer_finalization_failure_stays_in_worker_and_success_finalizes_once( - monkeypatch, -): - def make_transfer(finalize): - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer.backend = "test" - transfer.device = torch.device("cuda", 0) - transfer.global_rank = 0 - transfer.staging_depth = 1 - transfer.staging = [[]] - transfer._release = [threading.Event()] - transfer._release[0].set() - transfer._consumed_events = [object()] - transfer._consumed_recorded = [False] - transfer._changed_dst_slots = [()] - transfer._pending = deque() - transfer._pending_lock = threading.Lock() - transfer._error = None - transfer._thread = None - transfer._needs_staging_reuse_barrier = False - transfer._copy_batch = lambda _batch, _prepared_batch: None - transfer._start_transfer_generation = lambda: None - transfer._finish_transfer_generation = finalize - return transfer - - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - - finalized_before_publish = [] - success = make_transfer(lambda: finalized_before_publish.append(len(success._pending))) - success.start([(0, [])], [([(0, [], 0, [])], None)]) - success.finish() - assert finalized_before_publish == [0] - assert success.pending_layers() == [(0, 0)] - - failed = make_transfer(lambda: (_ for _ in ()).throw(RuntimeError("cache boom"))) - failed.start([(0, [])], [([(0, [], 0, [])], None)]) - failed._thread.join() - assert list(failed._pending) == [] - with pytest.raises(RuntimeError, match="EPLB migration worker failed") as exc_info: - failed.pending_layers() - assert isinstance(exc_info.value.__cause__, RuntimeError) - assert str(exc_info.value.__cause__) == "cache boom" - with pytest.raises(RuntimeError, match="EPLB migration worker failed"): - failed.finish() - assert failed._thread is None - - -def test_manager_rearms_after_rebalance_for_interval_one(): - recording_calls = [] - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.step_interval = 1 - manager.sampling_interval = 1 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = 0 - manager.weights = [] - manager._eplb_states = [] - manager.target_placement = torch.zeros(1) - manager.in_flight_started_at = 0 - manager.global_rank = 1 - manager._set_recording = lambda enabled: recording_calls.append(enabled) - manager._finish_rebalance() - assert manager.in_flight is False - assert recording_calls == [True] - assert manager._steady_collection_end_step is None - assert manager._continuous_collection_start_step is None - - -def test_manager_sparse_insufficient_schedules_bounded_fresh_window(monkeypatch): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "insufficient", - "minimum_layer_samples": 1, - "minimum": 2, - } - manager.global_rank = 0 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.prefill_steps = 37 - manager.weights = [] - manager._eplb_states = [] - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - recordings, resets, logs = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) - monkeypatch.setattr(manager_module.logger, "info", lambda *args: logs.append(args)) - - assert not manager._poll_evaluation() - assert not hasattr(manager, "_retained_local_samples") - assert manager._continuous_collection_start_step == 40 - assert manager._continuous_collection_end_step == 60 - assert recordings == [False] - assert resets == [True] - assert manager.sampling_interval == 20 - assert "insufficient samples" in logs[0][0] - assert "scheduled_fresh_window" in logs[0][0] - - manager.in_flight = False - manager.evaluation_in_flight = False - starts = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(manager.prefill_steps)) - manager.step() - manager.step() - assert manager.prefill_steps == 39 - assert starts == [] - manager.step() - assert manager.prefill_steps == 40 - assert starts == [] - for _ in range(19): - manager.step() - assert manager.prefill_steps == 59 - assert starts == [] - manager.step() - assert starts == [60] - - -def test_manager_full_window_insufficient_clears_and_backs_off(): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "insufficient", - "minimum_layer_samples": 1, - "minimum": 2, - } - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.prefill_steps = 60 - manager.weights = [] - manager._eplb_states = [] - manager._continuous_collection_start_step = 40 - manager._continuous_collection_end_step = 60 - recordings, resets = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) - - assert not manager._poll_evaluation() - assert not hasattr(manager, "_retained_local_samples") - assert manager._continuous_collection_start_step is None - assert manager._continuous_collection_end_step is None - assert manager.sampling_interval == 80 - assert recordings == [False] - assert resets == [True] - - -def test_begin_continuous_collection_uses_full_window_at_fixed_boundary(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = False - manager.prefill_steps = 36 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager._steady_collection_end_step = manager.prefill_steps + 1 - recordings, resets, starts = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) - - # The window is never truncated to the next boundary: it waits until 40, - # records a full 20 fresh steps, then evaluates at the boundary at 60. - manager._begin_continuous_collection() - assert manager._continuous_collection_start_step == 40 - assert manager._continuous_collection_end_step == 60 - assert recordings == [False] - assert resets == [True] - assert manager._steady_collection_end_step is None - - for _ in range(4): - manager.step() - assert manager.prefill_steps == 40 - assert recordings == [False, True] - assert starts == [] - for _ in range(19): - manager.step() - assert starts == [] - manager.step() # 60: the full window ends and triggers the evaluation. - assert starts == [True] - - -def test_begin_continuous_collection_preserves_full_window_at_sparse_boundary(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.prefill_steps = 80 - manager.step_interval = 20 - manager.sampling_interval = 80 - manager._steady_collection_end_step = None - recordings = [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: None - - manager._begin_continuous_collection() - assert manager._continuous_collection_start_step == 140 - assert manager._continuous_collection_end_step == 160 - assert recordings == [False] - - -def test_first_no_improvement_switches_to_sparse_sampling_window(monkeypatch): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager._continuous_collection_start_step = 0 - manager.weights = [] - manager._eplb_states = [] - recordings = [] - manager._set_recording = lambda enabled: recordings.append(enabled) - - assert not manager._poll_evaluation() - assert manager._continuous_collection_start_step is None - assert recordings == [False] - assert manager.sampling_interval == 80 - - -def test_continuous_collection_evaluates_only_after_one_full_base_window(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager._continuous_collection_start_step = 0 - manager._continuous_collection_end_step = 20 - manager.prefill_steps = 0 - manager.step_interval = 20 - manager.sampling_interval = 320 - manager._steady_collection_end_step = None - manager.evaluation_in_flight = False - started = [] - monkeypatch.setattr(manager, "_start_evaluation", lambda: started.append(True)) - - for _ in range(19): - manager.step() - assert manager.prefill_steps == 19 - assert started == [] - - manager.step() - assert manager.prefill_steps == 20 - assert started == [True] - - -def test_no_improvement_exponentially_backs_off_sampling_interval_at_cap(): - class DoneThread: - def join(self): - pass - - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager._evaluation_lock = threading.Lock() - manager._set_recording = lambda _enabled: None - manager.global_rank = 1 - manager.step_interval = 20 - manager.sampling_interval = 20 - manager.weights = [] - manager._eplb_states = [] - - for expected_interval in (80, 320, 320): - manager.evaluation_in_flight = True - manager._evaluation_error = None - manager._evaluation_thread = DoneThread() - manager._evaluation_result = { - "kind": "no_improvement", - "model_imbalance_ratio": 1.2, - "candidate_model_imbalance_ratio": 1.1, - "candidate_rebalance_gain": 0.01, - "candidate_changed_layer_count": 2, - } - assert not manager._poll_evaluation() - assert manager.sampling_interval == expected_interval - - -def test_sparse_backoff_arms_and_evaluates_only_at_new_interval_boundary(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.prefill_steps = 18 - manager.step_interval = 20 - manager.sampling_interval = 80 - manager._steady_collection_end_step = None - manager._continuous_collection_start_step = None - manager._continuous_collection_end_step = None - manager.evaluation_in_flight = False - recordings, resets, starts = [], [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._reset_recorded_samples = lambda: resets.append(True) - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) - - manager.step() - manager.step() - assert manager.prefill_steps == 20 - assert recordings == [] - assert starts == [] - - for _ in range(55): - manager.step() - assert manager.prefill_steps == 75 - assert recordings == [] - assert starts == [] - - manager.step() - assert manager.prefill_steps == 76 - assert recordings == [True] - assert resets == [True] - assert manager._steady_collection_end_step is not None - - for _ in range(4): - manager.step() - assert manager.prefill_steps == 80 - assert starts == [True] - assert manager._steady_collection_end_step is None - - -def test_planned_rebalance_resets_sampling_interval_to_base(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.current_placement = torch.zeros((1, 1, 1), dtype=torch.int64) - manager.num_logical_experts = 1 - manager.world_size = 1 - manager.node_world_size = 1 - manager.step_interval = 20 - manager.sampling_interval = 320 - manager._continuous_collection_start_step = 0 - manager.global_rank = 1 - manager.transfer = type( - "Transfer", - (), - {"start": lambda self, plans, prepared_batches: setattr(self, "started", (plans, prepared_batches))}, - )() - manager._reset_recorded_samples = lambda: None - - prepared_batches = [object()] - manager._start_rebalance( - { - "placement": torch.zeros((1, 1, 1), dtype=torch.int64), - "improved": torch.tensor([True]), - "metadata": [None], - "layer_plans": [(0, object())], - "prepared_batches": prepared_batches, - "before": {"max": 1.0, "p95": 1.0}, - "after": {"max": 1.0, "p95": 1.0}, - "model_imbalance_ratio": 1.0, - "candidate_model_imbalance_ratio": 1.0, - "candidate_rebalance_gain": 0.1, - "candidate_changed_layer_count": 1, - } - ) - - assert manager.sampling_interval == 20 - assert manager.in_flight - assert manager._continuous_collection_start_step is None - assert len(manager.transfer.started[0]) == 1 - assert manager.transfer.started[1] is prepared_batches - - -def test_transfer_start_rejects_prepared_batches_with_wrong_batch_count(): - transfer = object.__new__(transfer_module._EPLBTransferBase) - transfer._thread = None - transfer.staging_depth = 2 - transfer.staging = [object(), object()] - - with pytest.raises(ValueError, match="prepared batch count"): - transfer.start([(0, []), (1, []), (2, [])], prepared_batches=[object()]) - - -def test_first_rebalance_completion_switches_to_four_step_sparse_window(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.step_interval = 20 - manager.sampling_interval = manager.step_interval - manager._steady_collection_end_step = None - manager.weights = [] - manager._eplb_states = [] - manager.target_placement = torch.zeros(1) - manager.in_flight_started_at = 0 - manager.global_rank = 1 - recordings, starts = [], [] - manager._set_recording = lambda enabled: recordings.append(enabled) - manager._finish_rebalance() - assert recordings == [False] - - manager.in_flight = False - manager.prefill_steps = 38 - manager.evaluation_in_flight = False - monkeypatch.setattr(manager, "_start_evaluation", lambda: starts.append(True)) - manager.prefill_steps = 35 - manager.step() - assert recordings == [False, True] - for _ in range(4): - manager.step() - assert starts == [True] - - -def test_manager_inflight_step_does_not_poll(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = True - calls = [] - manager._poll_in_flight = lambda: calls.append("poll") - manager.step() - assert calls == [] - - -def test_manager_poll_advances_inflight_before_evaluation(): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = True - manager.evaluation_in_flight = True - calls = [] - manager._poll_in_flight = lambda: calls.append("inflight") - manager._poll_evaluation = lambda: calls.append("evaluation") - - manager.poll() - - assert calls == ["inflight"] - - -def test_manager_poll_waits_for_all_evaluation_results(monkeypatch): - manager = manager_module.EPLBManager.__new__(manager_module.EPLBManager) - manager.in_flight = False - manager.evaluation_in_flight = True - manager._evaluation_lock = threading.Lock() - manager._evaluation_result = None - manager._evaluation_error = None - manager.control_group = object() - manager._control_ready_count = torch.empty(1, dtype=torch.int32) - calls = [] - manager._poll_evaluation = lambda: calls.append("evaluation") - - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(0)) - manager.poll() - assert calls == [] - - manager._evaluation_result = {"kind": "no_improvement"} - monkeypatch.setattr(manager_module.dist, "all_reduce", lambda tensor, **_kwargs: tensor.fill_(1)) - manager.poll() - assert calls == ["evaluation"] - - -def test_nixl_descriptor_runs_merge_only_jointly_contiguous_source_and_destination_rows(): - steps = [ - TransferStep(0, 2, 1, 3), - TransferStep(0, 0, 1, 1), - TransferStep(0, 1, 1, 2), - TransferStep(0, 4, 1, 7), - ] - runs = transfer_module.NixlEPLBTransfer._contiguous_runs(steps) - assert [[(step.dst_slot, step.src_local_row) for step in run] for run in runs] == [ - [(0, 1), (1, 2), (2, 3)], - [(4, 7)], - ] - - -def test_nixl_prepare_batch_compiles_hot_path_without_tensor_views(monkeypatch): - class BatchMemcpy: - def __init__(self): - self.prepared = [] - self.enqueued = [] - - def prepare(self, descriptors): - descriptor = tuple(descriptors) - self.prepared.append(descriptor) - return descriptor - - def enqueue(self, descriptor, stream): - self.enqueued.append((descriptor, stream)) - - stream = SimpleNamespace(cuda_stream=123, synchronize=lambda: None) - batch_memcpy = BatchMemcpy() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = stream - transfer._batch_memcpy = batch_memcpy - transfer.global_rank = 0 - transfer._same_node_ranks = {0, 1, 2} - transfer._live_row_layout = [ - [ - ("w13.weight", 1000, 32), - ("w13.weight_scale", 2000, 32), - ("w2.weight", 3000, 32), - ] - ] - transfer._push_staging_row_layout = { - 1: [[("w13.weight", 4000, 32), ("w13.weight_scale", 5000, 32), ("w2.weight", 6000, 32)]], - 2: [[("w13.weight", 7000, 32), ("w13.weight_scale", 8000, 32), ("w2.weight", 9000, 32)]], - } - transfer._get_remote_read = lambda *_args: None - transfer._wait_xfers = lambda _xfers: None - - run_a = [TransferStep(1, 3, 0, 5), TransferStep(1, 4, 0, 6)] - run_b = [TransferStep(2, 1, 0, 2)] - local_inbound = TransferStep(0, 0, 1, 0) - remote_inbound = TransferStep(0, 1, 3, 2) - staging = object() - batch = [(0, run_a + run_b + [local_inbound, remote_inbound], 0, staging)] - prepared = transfer._prepare_batch(batch) - monkeypatch.setattr(transfer, "_prepare_batch", lambda _batch: pytest.fail("hot path must not prepare descriptors")) - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) - transfer._copy_batch(batch, prepared) - - expected = ( - (1160, 4096, 64), - (2160, 5096, 64), - (3160, 6096, 64), - (1064, 7032, 32), - (2064, 8032, 32), - (3064, 9032, 32), - ) - assert batch_memcpy.prepared == [expected] - assert batch_memcpy.enqueued == [(expected, 123)] - assert prepared.remote_entries == {3: [(0, [remote_inbound], staging)]} - - -def test_nixl_prepare_transfer_batches_match_staging_depth(): - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.staging_depth = 2 - transfer.staging = ["staging-0", "staging-1"] - seen_batches = [] - - def prepare_batch(batch): - seen_batches.append(batch) - return f"prepared-{len(seen_batches)}" - - transfer._prepare_batch = prepare_batch - layer_plans = [(3, "plan-3"), (4, "plan-4"), (5, "plan-5")] - - prepared_batches = transfer.prepare_transfer(layer_plans) - - assert prepared_batches == [ - (seen_batches[0], "prepared-1"), - (seen_batches[1], "prepared-2"), - ] - assert [[(layer, buffer) for layer, _plan, buffer, _staging in batch] for batch in seen_batches] == [ - [(3, 0), (4, 1)], - [(5, 0)], - ] - - -def test_cuda_batch_memcpy_cuda13_abi_and_descriptor_layout(): - class Function: - def __init__(self, callback): - self.callback = callback - self.restype = None - self.argtypes = None - - def __call__(self, *args): - return self.callback(*args) - - class Library: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *args: self.calls.append(args) or 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - self.calls = [] - - import ctypes - - library = Library() - batch_memcpy = _CudaBatchMemcpy(library) - prepared = batch_memcpy.prepare(((101, 201, 64), (102, 202, 128))) - batch_memcpy.enqueue(prepared, 777) - - assert len(library.cudaMemcpyBatchAsync.argtypes) == 8 - assert library.calls[0][3] == 2 - assert library.calls[0][6] == 1 - assert library.calls[0][7].value == 777 - assert [pointer for pointer in library.calls[0][0]] == [201, 202] - assert [pointer for pointer in library.calls[0][1]] == [101, 102] - assert list(library.calls[0][2]) == [64, 128] - attrs = library.calls[0][4]._obj - assert attrs.srcAccessOrder == 1 - assert attrs.srcLocHint.type == attrs.srcLocHint.id == 0 - assert attrs.dstLocHint.type == attrs.dstLocHint.id == 0 - assert attrs.flags == 1 - - -def test_cuda_batch_memcpy_rejects_unsupported_runtime_and_invalid_descriptors(): - class Function: - def __init__(self, callback): - self.callback = callback - self.restype = None - self.argtypes = None - - def __call__(self, *args): - return self.callback(*args) - - class OldRuntimeLibrary: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 12080 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - - import ctypes - - with pytest.raises(RuntimeError, match="13.0"): - _CudaBatchMemcpy(OldRuntimeLibrary()) - - class FutureRuntimeLibrary: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 14000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaMemcpyBatchAsync = Function(lambda *_args: 0) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - - with pytest.raises(RuntimeError, match="13.x"): - _CudaBatchMemcpy(FutureRuntimeLibrary()) - - class MissingBatchSymbolLibrary: - def __init__(self): - def get_version(pointer): - ctypes.cast(pointer, ctypes.POINTER(ctypes.c_int))[0] = 13000 - return 0 - - self.cudaRuntimeGetVersion = Function(get_version) - self.cudaGetErrorString = Function(lambda _result: b"fake cuda error") - - with pytest.raises(RuntimeError, match="cudaMemcpyBatchAsync"): - _CudaBatchMemcpy(MissingBatchSymbolLibrary()) - with pytest.raises(ValueError, match="at least one"): - _CudaBatchMemcpy.prepare(()) - with pytest.raises(ValueError, match="positive"): - _CudaBatchMemcpy.prepare(((1, 2, 0),)) - - -def test_nixl_transfer_fails_fast_without_cuda13_batch_memcpy(monkeypatch): - failure = RuntimeError("missing cudaMemcpyBatchAsync") - - def unavailable(): - raise failure - - monkeypatch.setattr( - transfer_module._EPLBTransferBase, "__init__", lambda self, *_args: setattr(self, "device", "mock") - ) - monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) - monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", unavailable) - with pytest.raises(RuntimeError, match="missing cudaMemcpyBatchAsync") as exc_info: - transfer_module.NixlEPLBTransfer([object()], object(), 0, 1) - assert exc_info.value is failure - - -@pytest.mark.parametrize( - ("maps", "open_raises", "expected"), - [ - ( - "7f /cuda-12/libcudart.so.12.8\n" - "7f /first/libcudart.so.13 (deleted)\n" - "7f /second/libcudart.so.13\n" - "7f /cuda-14/libcudart.so.14\n" - "7f /cuda-130/libcudart.so.130\n", - False, - "/first/libcudart.so.13", - ), - ("7f /cuda-12/libcudart.so.12\n7f /cuda-130/libcudart.so.130\n", False, None), - ("", True, None), - ], -) -def test_cuda_batch_memcpy_finds_first_loaded_cuda13_runtime(monkeypatch, maps, open_raises, expected): - def open_maps(_path): - if open_raises: - raise OSError("maps unavailable") - return io.StringIO(maps) - - monkeypatch.setattr(builtins, "open", open_maps) - assert _CudaBatchMemcpy._find_loaded_cudart() == expected - - -@pytest.mark.parametrize(("layers", "expected_depth"), [(1, 1), (8, 8), (9, 8), (43, 8)]) -def test_nixl_transfer_bounds_staging_depth(monkeypatch, layers, expected_depth): - def base_init(self, *_args): - self.device = "mock-device" - - monkeypatch.setattr(transfer_module._EPLBTransferBase, "__init__", base_init) - monkeypatch.setattr(transfer_module.torch.cuda, "Stream", lambda device: ("stream", device)) - monkeypatch.setattr(transfer_module, "_CudaBatchMemcpy", lambda: object()) - monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_ipc_metadata", lambda self: None) - monkeypatch.setattr(transfer_module.NixlEPLBTransfer, "_init_push_layouts", lambda self: None) - - transfer = transfer_module.NixlEPLBTransfer([object()] * layers, object(), 0, 1) - - assert transfer.staging_depth == expected_depth - - -def test_nixl_remote_read_cache_reuses_exact_batch_key_and_releases_on_shutdown(): - class Tensor: - nbytes = 8 - - def __init__(self, pointer): - self.pointer = pointer - - def data_ptr(self): - return self.pointer - - def get_device(self): - return 0 - - def __getitem__(self, _): - return self - - class Agent: - def __init__(self): - self.prepared = 0 - self.made = 0 - self.released_xfers = 0 - self.released_dlists = 0 - self.removed_agents = [] - - def get_xfer_descs(self, descriptors, _): - return descriptors - - def prep_xfer_dlist(self, *_args, **_kwargs): - self.prepared += 1 - return f"dlist-{self.prepared}" - - def make_prepped_xfer(self, *_args, **_kwargs): - self.made += 1 - return f"xfer-{self.made}" - - def query_xfer_backend(self, _): - return "UCX" - - def release_xfer_handle(self, _): - self.released_xfers += 1 - - def release_dlist_handle(self, _): - self.released_dlists += 1 - - def remove_remote_agent(self, remote_name): - self.removed_agents.append(remote_name) - - agent = Agent() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._nixl_agent = agent - transfer._xfer_cache = {} - transfer._used_xfer_cache_keys = set() - transfer._remote_agents = {1: "remote-1"} - transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} - transfer._registered_descs = None - transfer.live = [[("weight", Tensor(100))]] - staging = [("weight", Tensor(200))] - first_entries = [(0, [TransferStep(0, 0, 1, 0)], staging)] - - first = transfer._get_remote_read(1, first_entries) - assert transfer._get_remote_read(1, first_entries) is first - assert agent.made == 1 - - changed_entries = [(0, [TransferStep(0, 1, 1, 0)], staging)] - transfer._get_remote_read(1, changed_entries) - assert agent.made == 2 - - transfer.shutdown() - assert agent.released_xfers == 2 - assert agent.released_dlists == 4 - assert agent.removed_agents == ["remote-1"] - - -def test_nixl_remote_read_cache_is_bounded_to_the_current_transfer_generation( - monkeypatch, -): - class Tensor: - nbytes = 8 - - def __init__(self, pointer): - self.pointer = pointer - - def data_ptr(self): - return self.pointer - - def get_device(self): - return 0 - - def __getitem__(self, _): - return self - - class Agent: - def __init__(self): - self.made = 0 - self.released_xfers = 0 - self.released_dlists = 0 - - def get_xfer_descs(self, descriptors, _): - return descriptors - - def prep_xfer_dlist(self, *_args, **_kwargs): - return object() - - def make_prepped_xfer(self, *_args, **_kwargs): - self.made += 1 - return object() - - def query_xfer_backend(self, _): - return "UCX" - - def release_xfer_handle(self, _): - self.released_xfers += 1 - - def release_dlist_handle(self, _): - self.released_dlists += 1 - - def remove_remote_agent(self, _): - pass - - agent = Agent() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._nixl_agent = agent - transfer._xfer_cache = {} - transfer._used_xfer_cache_keys = set() - transfer._remote_agents = {1: "remote-1"} - transfer._remote_layouts = {1: [[("weight", 1000, 0, 8)]]} - transfer._registered_descs = None - transfer._ipc_staging = {} - transfer.live = [[("weight", Tensor(100))]] - transfer.device = torch.device("cuda", 0) - transfer.transfer_group = object() - transfer.global_rank = 0 - transfer.world_size = 2 - transfer.staging_depth = 1 - transfer.staging = [[]] - transfer.num_experts_per_rank = 0 - transfer._release = [threading.Event()] - transfer._release[0].set() - transfer._consumed_events = [object()] - transfer._consumed_recorded = [False] - transfer._changed_dst_slots = [()] - transfer._pending = deque() - transfer._pending_lock = threading.Lock() - transfer._error = None - transfer._thread = None - transfer._needs_staging_reuse_barrier = False - - staging = [("weight", Tensor(200))] - entries_a = [(0, [TransferStep(0, 0, 1, 0)], staging)] - entries_b = [(0, [TransferStep(0, 1, 1, 0)], staging)] - generation = [entries_a] - - def copy_batch(_batch, _prepared_batch): - transfer._get_remote_read(1, generation[0]) - - transfer._copy_batch = copy_batch - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - - prepared_batches = [([(0, [], 0, transfer.staging[0])], None)] - transfer.start([(0, [])], prepared_batches) - transfer.finish() - assert agent.made == 1 - assert agent.released_xfers == agent.released_dlists == 0 - assert len(transfer._xfer_cache) == 1 - - # The real manager releases each staging buffer through commit(). This - # focused cache test has no commits, so model that hand-off before the - # next generation reuses buffer zero. - transfer._release[0].set() - transfer.start([(0, [])], prepared_batches) - transfer.finish() - assert agent.made == 1 - assert agent.released_xfers == agent.released_dlists == 0 - assert len(transfer._xfer_cache) == 1 - - generation[:] = [entries_b] - transfer._release[0].set() - transfer.start([(0, [])], prepared_batches) - transfer.finish() - assert agent.made == 2 - assert agent.released_xfers == 1 - assert agent.released_dlists == 2 - assert len(transfer._xfer_cache) == 1 - transfer.shutdown() - - -def test_nixl_ipc_metadata_exports_staging_per_local_target(monkeypatch): - from lightllm.server.router.model_infer.mode_backend.pd import p2p_fix - - class Tensor: - shape = (4, 2) - dtype = torch.float16 - device = torch.device("cuda", 0) - nbytes = 16 - - def __init__(self, label): - self.label = label - - def numel(self): - return 3 - - def __getitem__(self, _index): - return self - - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer.transfer_group = object() - transfer.global_rank = 0 - transfer.world_size = 3 - transfer.device = torch.device("cuda", 0) - transfer.live = [[("w13.weight", Tensor("local-w13")), ("w2.weight", Tensor("local-w2"))]] - transfer.staging_depth = 8 - transfer.staging = [[("w13.weight", Tensor(f"local-staging-{index}"))] for index in range(8)] - transfer._ipc_staging = {} - - reduce_calls, rebuild_calls, gathers = [], [], [] - - def reduce_tensor(tensor): - reduce_calls.append(tensor.label) - return None, (f"export-{tensor.label}",) - - def rebuild_tensor(export): - rebuild_calls.append(export) - return Tensor(export) - - source_one = [[("w13.weight", (4, 2), torch.float16, (f"rank1-staging-{index}",))] for index in range(8)] - - def all_gather(output, value, **_kwargs): - gathers.append(value) - if len(gathers) == 1: - output[:] = ["node-a", "node-a", "node-b"] - else: - output[:] = [value, {0: {"staging": source_one}}, {}] - - monkeypatch.setattr(p2p_fix, "reduce_tensor", reduce_tensor) - monkeypatch.setattr(p2p_fix, "p2p_fix_rebuild_cuda_tensor", rebuild_tensor) - monkeypatch.setattr(transfer_module.dist, "all_gather_object", all_gather) - monkeypatch.setattr(transfer_module.torch.cuda, "set_device", lambda _device: None) - - transfer._init_ipc_metadata() - - assert reduce_calls == [*(f"local-staging-{index}" for index in range(8))] - assert rebuild_calls == [*(f"rank1-staging-{index}" for index in range(8))] - assert transfer._same_node_ranks == {0, 1} - assert transfer._cross_node_ranks == {2} - assert transfer._needs_staging_reuse_barrier - assert [name for name, _ in transfer._ipc_staging[1][0]] == ["w13.weight"] - - -def test_nixl_copy_batch_source_pushes_local_rows_and_keeps_remote_ucx_reads( - monkeypatch, -): - class Stream: - def __init__(self): - self.synchronized = 0 - self.cuda_stream = 123 - - def synchronize(self): - self.synchronized += 1 - - stream = Stream() - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = stream - transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) - remote_reads, waited_xfers, enqueued = [], [], [] - transfer._get_remote_read = lambda src_rank, entries: remote_reads.append((src_rank, entries)) or ( - None, - None, - "xfer", - ) - transfer._wait_xfers = waited_xfers.extend - - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) - prepared_push = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") - transfer._copy_batch([], prepared_push) - assert enqueued == [("push", stream.cuda_stream)] - - remote_step = TransferStep(0, 1, 2, 0) - prepared_remote = transfer_module.NixlEPLBTransfer._PreparedBatch({2: [(0, [remote_step], [])]}, None) - transfer._copy_batch([], prepared_remote) - assert [rank for rank, _ in remote_reads] == [2] - assert waited_xfers == [(None, None, "xfer")] - - -def test_nixl_copy_batch_self_only_rank_pushes_and_synchronizes(monkeypatch): - class Stream: - def __init__(self): - self.synchronized = 0 - self.cuda_stream = 456 - - def synchronize(self): - self.synchronized += 1 - - transfer = object.__new__(transfer_module.NixlEPLBTransfer) - transfer._push_stream = Stream() - enqueued = [] - transfer._batch_memcpy = SimpleNamespace(enqueue=lambda descriptor, stream: enqueued.append((descriptor, stream))) - transfer._wait_xfers = lambda _xfers: None - monkeypatch.setattr(transfer_module.torch.cuda, "stream", lambda _stream: pytest.fail("must not switch streams")) - - prepared = transfer_module.NixlEPLBTransfer._PreparedBatch({}, "push") - transfer._copy_batch([], prepared) - - assert enqueued == [("push", transfer._push_stream.cuda_stream)] - assert transfer._push_stream.synchronized == 1 - - -def test_manager_constructs_nixl_transfer(monkeypatch): - weight = type( - "Weight", - (), - { - "n_routed_experts": 4, - "expert_parallel_state": _test_parallel_state( - eplb=True, - num_logical_experts=4, - world_size=2, - num_redundant_experts_per_rank=2, - route_counter=torch.zeros((2, 4), dtype=torch.int64), - ), - }, - )() - transfer = object() - groups = [object(), object(), object()] - new_group_calls = [] - monkeypatch.setattr(manager_module, "_find_fused_moe_weights", lambda model: [weight]) - monkeypatch.setattr(manager_module, "get_global_rank", lambda: 0) - monkeypatch.setattr(manager_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(manager_module, "get_node_world_size", lambda: 2) - monkeypatch.setattr(manager_module, "get_prefill_eplb_step_interval", lambda: 20) - monkeypatch.setattr(manager_module, "get_eplb_rebalance_gain_threshold", lambda: 0.07) - - def new_group(*args, **kwargs): - new_group_calls.append((args, kwargs)) - return groups[len(new_group_calls) - 1] - - monkeypatch.setattr(manager_module.dist, "new_group", new_group) - transfer_calls = [] - monkeypatch.setattr( - manager_module, - "NixlEPLBTransfer", - lambda weights, group, rank, world_size: ( - transfer_calls.append((weights, group, rank, world_size)) or transfer - ), - ) - logs = [] - monkeypatch.setattr(manager_module.logger, "info", lambda message: logs.append(message)) - manager = manager_module.EPLBManager(type("Model", (), {})()) - assert manager.transfer is transfer - assert ( - manager.evaluation_group, - manager.control_group, - manager.transfer_group, - ) == tuple(groups) - assert new_group_calls == [(([0, 1],), {"backend": "gloo"})] * 3 - assert transfer_calls == [([weight], groups[2], 0, 2)] - assert manager.rebalance_gain_threshold == 0.07 - assert "rebalance_gain_threshold=0.0700" in logs[0] - assert manager._continuous_collection_start_step is None - assert manager._continuous_collection_end_step == manager.step_interval - assert weight.expert_parallel_state.eplb.recording - assert manager._eplb_states[0] is weight.expert_parallel_state.eplb - assert not hasattr(weight.expert_parallel_state.eplb, "record_load") - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -@pytest.mark.parametrize("record_load", [False, True]) -@pytest.mark.parametrize("tokens", [1, 32]) -@pytest.mark.parametrize("scoring_func", ["sigmoid", "softmax"]) -@pytest.mark.parametrize("renormalize", [False, True]) -def test_grouped_topk_eplb_matches_topk_mapping_and_counting(record_load, tokens, scoring_func, renormalize): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import ( - triton_grouped_topk, - triton_grouped_topk_eplb, - ) - - torch.manual_seed(1234) - topk = 8 - experts = 256 - num_expert_group = 8 - gating_output = torch.randn((tokens, experts), dtype=torch.bfloat16, device="cuda") - correction_bias = torch.randn((experts,), dtype=torch.float32, device="cuda") - hidden_states = torch.empty((tokens, 1), dtype=torch.bfloat16, device="cuda") - logical_to_physical = torch.stack( - ( - torch.arange(experts, dtype=torch.int32, device="cuda"), - torch.arange(experts, dtype=torch.int32, device="cuda") + experts, - ), - dim=1, - ) - logical_replica_count = torch.where( - torch.arange(experts, device="cuda") % 3 == 0, - torch.full((experts,), 2, dtype=torch.int32, device="cuda"), - torch.ones((experts,), dtype=torch.int32, device="cuda"), - ) - expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") - fused_counter = torch.zeros_like(expected_counter) - - expected_weights, logical_ids = triton_grouped_topk( - hidden_states, - gating_output, - correction_bias, - topk, - renormalize, - num_expert_group, - 4, - scoring_func, - 2, - ) - if tokens == 1: - replica_indices = torch.zeros_like(logical_ids) - else: - token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) - replica_indices = ( - (((token_indices * 2654435769) & 0xFFFFFFFF) + ((logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF)) - & 0xFFFFFFFF - ) % logical_replica_count[logical_ids.to(torch.long)].to(torch.int64) - expected_ids = logical_to_physical[logical_ids.to(torch.long), replica_indices.to(torch.long)].to(torch.long) - if record_load: - expected_counter[1].scatter_add_( - 0, - logical_ids.reshape(-1).to(torch.long), - torch.ones(logical_ids.numel(), dtype=torch.int64, device="cuda"), - ) - fused_weights, fused_ids, fused_logical_ids = triton_grouped_topk_eplb( - gating_output, - correction_bias, - topk, - renormalize, - num_expert_group, - 4, - scoring_func, - logical_to_physical, - logical_replica_count, - fused_counter, - sample_index=1, - record_load=record_load, - use_grouped_topk=True, - group_score_used_topk_num=2, - ) - torch.cuda.synchronize() - - torch.testing.assert_close(fused_weights, expected_weights, rtol=1e-5, atol=1e-6) - assert torch.equal(fused_ids, expected_ids) - assert fused_logical_ids is None - assert torch.equal(fused_counter, expected_counter) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -@pytest.mark.parametrize("record_load", [False, True]) -@pytest.mark.parametrize("tokens", [1, 32]) -def test_global_topk_eplb_supports_logical_ids_and_counting(record_load, tokens): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - - torch.manual_seed(1234) - topk = 4 - experts = 64 - gating_output = torch.randn((tokens, experts), dtype=torch.float32, device="cuda") - logical_to_physical = torch.stack( - ( - torch.arange(experts, dtype=torch.int32, device="cuda"), - torch.arange(experts, dtype=torch.int32, device="cuda") + experts, - ), - dim=1, - ) - logical_replica_count = torch.where( - torch.arange(experts, device="cuda") % 3 == 0, - torch.full((experts,), 2, dtype=torch.int32, device="cuda"), - torch.ones((experts,), dtype=torch.int32, device="cuda"), - ) - expected_counter = torch.zeros((2, experts), dtype=torch.int64, device="cuda") - fused_counter = torch.zeros_like(expected_counter) - expected_weights, expected_logical_ids = torch.softmax(gating_output, dim=-1).topk(topk, dim=-1) - if tokens == 1: - replica_indices = torch.zeros_like(expected_logical_ids) - else: - token_indices = torch.arange(tokens, device="cuda", dtype=torch.int64).unsqueeze(1) - replica_indices = ( - ( - ((token_indices * 2654435769) & 0xFFFFFFFF) - + ((expected_logical_ids.to(torch.int64) * 2246822519) & 0xFFFFFFFF) - ) - & 0xFFFFFFFF - ) % logical_replica_count[expected_logical_ids].to(torch.int64) - expected_ids = logical_to_physical[expected_logical_ids, replica_indices.to(torch.long)].to(torch.long) - if record_load: - expected_counter[1].scatter_add_( - 0, - expected_logical_ids.reshape(-1), - torch.ones(expected_logical_ids.numel(), dtype=torch.int64, device="cuda"), - ) - expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) - - weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - gating_output=gating_output, - correction_bias=torch.randn((experts,), dtype=torch.float32, device="cuda"), - topk=topk, - renormalize=True, - num_expert_group=8, - topk_group=4, - scoring_func="sigmoid", - logical_to_physical_map=logical_to_physical, - logical_replica_count=logical_replica_count, - expert_counter=fused_counter, - sample_index=1, - record_load=record_load, - use_grouped_topk=False, - return_logical_ids=True, - ) - torch.cuda.synchronize() - - torch.testing.assert_close(weights, expected_weights, rtol=1e-5, atol=1e-6) - assert torch.equal(physical_ids, expected_ids) - assert torch.equal(logical_ids, expected_logical_ids) - assert torch.equal(fused_counter, expected_counter) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required for the Triton EPLB kernel") -def test_triton_grouped_topk_eplb_empty_tokens_skips_kernel(): - from lightllm.common.basemodel.triton_kernel.fused_moe.grouped_topk import triton_grouped_topk_eplb - - experts = 64 - counter = torch.zeros((1, experts), dtype=torch.int64, device="cuda") - weights, physical_ids, logical_ids = triton_grouped_topk_eplb( - gating_output=torch.empty((0, experts), device="cuda"), - correction_bias=None, - topk=4, - renormalize=False, - num_expert_group=8, - topk_group=4, - scoring_func="softmax", - logical_to_physical_map=torch.zeros((experts, 1), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones((experts,), dtype=torch.int32, device="cuda"), - expert_counter=counter, - sample_index=0, - record_load=True, - use_grouped_topk=False, - return_logical_ids=True, - ) - - assert weights.shape == physical_ids.shape == logical_ids.shape == (0, 4) - assert physical_ids.dtype is logical_ids.dtype is torch.long - assert torch.equal(counter, torch.zeros_like(counter)) diff --git a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py b/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py deleted file mode 100644 index 3eff70f940..0000000000 --- a/unit_tests/common/fused_moe/test_eplb_transfer_gpu.py +++ /dev/null @@ -1,444 +0,0 @@ -"""NIXL EPLB correctness tests and a two-GPU 512 MiB micro-performance test.""" -import os -import random -import socket -import statistics -import time - -import pytest -import torch -import torch.distributed as dist -import torch.multiprocessing as mp - -from lightllm.server.router.model_infer.mode_backend.eplb_transfer import ( - NixlEPLBTransfer, - align_target_placement, - build_transfer_plan, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.eplb_placement import ( - build_initial_redundant_expert_ids, -) -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.expert_parallel_state import ( - EPLBState, - ExpertParallelState, -) - -pytest.importorskip("nixl", reason="NIXL package is required") - - -class _Pack: - def __init__(self, weight, weight_scale): - self.weight = weight - self.weight_scale = weight_scale - self.weight_zero_point = None - - -def _free_port(): - sock = socket.socket() - sock.bind(("127.0.0.1", 0)) - port = sock.getsockname()[1] - sock.close() - return port - - -class _FakeWeight: - def __init__(self, rank, layer_index, row_elements): - self.n_routed_experts = 32 - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=32, - world_size=2, - eplb=EPLBState( - num_redundant_experts_per_rank=16, - initial_redundant_expert_ids_by_rank=build_initial_redundant_expert_ids(32, 2, 16), - logical_to_physical_map=torch.zeros((32, 2), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones(32, dtype=torch.int32, device="cuda"), - route_counter=torch.zeros((1, 32), dtype=torch.int64, device="cuda"), - ), - ) - base = rank * 100 + layer_index * 100 - self.w13 = self._pack(base, row_elements) - self.w2 = self._pack(base + 10, row_elements) - - @staticmethod - def _pack(base, row_elements): - weight = torch.empty((32, row_elements), dtype=torch.float16, device="cuda") - for row in range(weight.shape[0]): - weight[row].fill_(base + row) - scale = torch.empty((32, 1), dtype=torch.float32, device="cuda") - for row in range(scale.shape[0]): - scale[row].fill_(base + row + 0.5) - return _Pack(weight, scale) - - -def _wait_for_ready_prefix(transfer, control_group): - deadline = time.monotonic() + 30 - while True: - pending = transfer.pending_layers() - ready_count = torch.tensor([len(pending)], dtype=torch.int32) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=control_group) - if int(ready_count.item()) > 0: - return pending[: int(ready_count.item())] - if time.monotonic() >= deadline: - raise TimeoutError("EPLB transfer worker did not publish a globally ready layer") - time.sleep(0.001) - - -def _run_layers( - transfer, - control_group, - layer_plans, - callback=lambda layer_index: None, - before_commit_callback=lambda layer_index: None, -): - transfer.start(layer_plans, transfer.prepare_transfer(layer_plans)) - committed = 0 - while committed < len(layer_plans): - pending = _wait_for_ready_prefix(transfer, control_group) - assert len(pending) <= len(layer_plans) - committed - for layer_index, buffer_index in pending: - assert layer_index == layer_plans[committed][0] - before_commit_callback(layer_index) - transfer.commit( - layer_index, - buffer_index, - lambda layer_index=layer_index: callback(layer_index), - ) - committed += 1 - transfer.finish() - - -def _assert_correctness(weights, rank, source_rows_by_dst_slot): - if rank == 0: - for layer_index in (0, len(weights) - 1): - base = 100 + layer_index * 100 - for dst_slot, src_row in enumerate(source_rows_by_dst_slot): - dst_row = 16 + dst_slot - assert torch.all(weights[layer_index].w13.weight[dst_row] == base + src_row) - assert torch.all(weights[layer_index].w13.weight_scale[dst_row] == base + src_row + 0.5) - assert torch.all(weights[layer_index].w2.weight[dst_row] == base + src_row + 10) - assert torch.all(weights[layer_index].w2.weight_scale[dst_row] == base + src_row + 10.5) - - -def _benchmark(transfer, control_group, layer_plans, payload): - for _ in range(3): - _run_layers(transfer, control_group, layer_plans) - dist.barrier(group=control_group) - started = time.perf_counter() - for _ in range(8): - _run_layers(transfer, control_group, layer_plans) - torch.cuda.synchronize() - return payload * 8 / (time.perf_counter() - started) / 1e9 - - -def _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload): - batch = [ - (layer_index, plan, buffer_index, transfer.staging[buffer_index]) - for buffer_index, (layer_index, plan) in enumerate(layer_plans) - ] - # Measure the precompiled hot path only; planning and descriptor construction are off the timer. - prepared_batch = transfer._prepare_batch(batch) - for _ in range(3): - transfer._copy_batch(batch, prepared_batch) - dist.barrier(group=control_group) - samples = [] - for _ in range(20): - started = time.perf_counter() - transfer._copy_batch(batch, prepared_batch) - samples.append(payload / (time.perf_counter() - started) / 1e9) - torch.cuda.synchronize() - median = statistics.median(samples) - print( - f"NIXL _copy_batch payload={payload / 2**20:.1f} MiB; " - f"min={min(samples):.2f} GB/s median={median:.2f} " - f"mean={statistics.mean(samples):.2f} max={max(samples):.2f}", - flush=True, - ) - return median - - -def _eplb_worker(rank, port, queue): - os.environ["MASTER_ADDR"] = "127.0.0.1" - os.environ["MASTER_PORT"] = str(port) - torch.cuda.set_device(rank) - dist.init_process_group("gloo", rank=rank, world_size=2) - control_group = dist.new_group([0, 1], backend="gloo") - transfer_group = dist.new_group([0, 1], backend="gloo") - # Eight layers × 16 changed experts × two 2 MiB rows = 512 MiB useful remote weight payload. - row_elements = int(os.getenv("LIGHTLLM_EPLB_TEST_ROW_ELEMENTS", str(1024 * 1024))) - current = torch.tensor([list(range(16)), list(range(16))]) - source_rows_by_dst_slot = list(range(16)) - random.Random(20260731).shuffle(source_rows_by_dst_slot) - # Rank 0 receives the reverse logical-expert range in a deterministic random slot order. - # Consequently every descriptor has a distinct source and destination row. - target = torch.tensor([[16 + source_row for source_row in source_rows_by_dst_slot], list(range(16))]) - plan = build_transfer_plan(current, target, 32, 2, 2) - assert [step.src_local_row for step in plan if step.dst_rank == 0] == source_rows_by_dst_slot - benchmark_layer_count = 8 - layer_count = benchmark_layer_count + 1 - row_payload = ( - 2 * row_elements * torch.empty((), dtype=torch.float16).element_size() - + 2 * torch.empty((), dtype=torch.float32).element_size() - ) - payload = benchmark_layer_count * 16 * row_payload - weights = [_FakeWeight(rank, layer_index, row_elements) for layer_index in range(layer_count)] - transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=2) - assert transfer.staging_depth == 8 - assert transfer._eplb_states[0] is weights[0].expert_parallel_state.eplb - assert all(tensor.is_cuda for staging in transfer.staging for _, tensor in staging) - - wrap_layer_plans = [(layer_index, plan) for layer_index in range(layer_count)] - - def delay_rank_zero_first_commit(layer_index): - if rank == 0 and layer_index == 0: - time.sleep(0.1) - - _run_layers( - transfer, - control_group, - wrap_layer_plans, - before_commit_callback=delay_rank_zero_first_commit, - ) - torch.cuda.synchronize() - _assert_correctness(weights, rank, source_rows_by_dst_slot) - layer_plans = wrap_layer_plans[:benchmark_layer_count] - nixl_copy_batch = _benchmark_nixl_copy_batch(transfer, control_group, layer_plans, payload) - nixl_bandwidth = _benchmark(transfer, control_group, layer_plans, payload) - transfer.shutdown() - - gathered = [None, None] - dist.all_gather_object(gathered, (nixl_bandwidth, nixl_copy_batch), group=control_group) - if rank == 0: - queue.put((payload, *gathered[0])) - dist.barrier(group=control_group) - dist.destroy_process_group() - - -@pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < 2, - reason="requires two CUDA GPUs", -) -def test_eplb_transfer_two_gpu_correctness_and_microperf(): - queue = mp.get_context("spawn").SimpleQueue() - mp.spawn(_eplb_worker, args=(_free_port(), queue), nprocs=2, join=True) - payload, nixl_gbps, nixl_copy_batch_gbps = queue.get() - print( - f"EPLB remote payload/round: {payload / 2**20:.1f} MiB; " - f"NIXL={nixl_gbps:.2f} GB/s NIXL _copy_batch={nixl_copy_batch_gbps:.2f} GB/s" - ) - assert nixl_gbps > 0 - - -def _depth_value(expert, layer_index, offset): - return expert // 32 * 100 + layer_index * 5 + expert % 32 + offset - - -class _DepthWeight: - def __init__(self, rank, layer_index, initial_placement): - self.n_routed_experts = 256 - self.expert_parallel_state = ExpertParallelState( - num_logical_experts=256, - world_size=8, - eplb=EPLBState( - num_redundant_experts_per_rank=4, - initial_redundant_expert_ids_by_rank=initial_placement.clone(), - logical_to_physical_map=torch.zeros((256, 8), dtype=torch.int32, device="cuda"), - logical_replica_count=torch.ones(256, dtype=torch.int32, device="cuda"), - route_counter=torch.zeros((1, 256), dtype=torch.int64, device="cuda"), - ), - ) - logical_ids = list(range(rank * 32, (rank + 1) * 32)) + initial_placement[rank].tolist() - self.w13 = self._pack(logical_ids, layer_index, 0) - self.w2 = self._pack(logical_ids, layer_index, 2) - - @staticmethod - def _pack(logical_ids, layer_index, offset): - weight = torch.empty((36, 64), dtype=torch.float16, device="cuda") - scale = torch.empty((36, 1), dtype=torch.float32, device="cuda") - for row, expert in enumerate(logical_ids): - value = _depth_value(expert, layer_index, offset) - weight[row].fill_(value) - scale[row].fill_(value + 0.25) - return _Pack(weight, scale) - - -def _depth_target(layer_index): - return torch.tensor([[((dst + layer_index + slot + 1) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]) - - -def _wait_all_pending(transfer, group, expected_count): - deadline = time.monotonic() + 30 - while True: - pending = transfer.pending_layers() - ready_count = torch.tensor([len(pending)], dtype=torch.int32) - dist.all_reduce(ready_count, op=dist.ReduceOp.MIN, group=group) - if int(ready_count.item()) == expected_count: - return pending - if time.monotonic() > deadline: - raise TimeoutError(f"expected {expected_count} pending layers, got {pending}") - time.sleep(0.001) - - -def _clone_depth_live(weights): - return [ - [tensor.detach().clone() for _, tensor in transfer_tensors] - for transfer_tensors in [ - [ - ("w13.weight", weight.w13.weight), - ("w13.scale", weight.w13.weight_scale), - ("w2.weight", weight.w2.weight), - ("w2.scale", weight.w2.weight_scale), - ] - for weight in weights - ] - ] - - -def _assert_depth_snapshot(weights, snapshot, layer_indices=None, primary_only=False): - if layer_indices is None: - layer_indices = range(len(weights)) - for layer_index in layer_indices: - live_tensors = ( - weights[layer_index].w13.weight, - weights[layer_index].w13.weight_scale, - weights[layer_index].w2.weight, - weights[layer_index].w2.weight_scale, - ) - for live, expected in zip(live_tensors, snapshot[layer_index]): - if primary_only: - live = live[:32] - expected = expected[:32] - torch.testing.assert_close(live, expected) - - -def _assert_depth_staging(rank, layer_plans, target_placements, pending, transfer): - for buffer_index, ((layer_index, plan), pending_item) in enumerate(zip(layer_plans, pending)): - assert pending_item == (layer_index, buffer_index) - expected = {step.dst_slot: step for step in plan if step.dst_rank == rank} - for dst_slot, step in expected.items(): - expert = int(target_placements[layer_index][rank, dst_slot]) - base = _depth_value(expert, layer_index, 0) - staging = transfer.staging[buffer_index] - assert torch.all(staging[0][1][dst_slot] == base) - assert torch.all(staging[1][1][dst_slot] == base + 0.25) - assert torch.all(staging[2][1][dst_slot] == base + 2) - assert torch.all(staging[3][1][dst_slot] == base + 2.25) - - -def _assert_depth_live(weights, rank, layer_plans, target_placements): - for layer_index, plan in layer_plans: - for step in plan: - if step.dst_rank != rank: - continue - expert = int(target_placements[layer_index][rank, step.dst_slot]) - base = _depth_value(expert, layer_index, 0) - assert torch.all(weights[layer_index].w13.weight[32 + step.dst_slot] == base) - assert torch.all(weights[layer_index].w13.weight_scale[32 + step.dst_slot] == base + 0.25) - assert torch.all(weights[layer_index].w2.weight[32 + step.dst_slot] == base + 2) - assert torch.all(weights[layer_index].w2.weight_scale[32 + step.dst_slot] == base + 2.25) - - -def _assert_peer_coverage(layer_plans, require_redundant_source=False): - steps = [step for _, plan in layer_plans for step in plan] - assert {step.dst_rank for step in steps} == set(range(8)) - assert {step.src_rank for step in steps} == set(range(8)) - for source_rank in range(8): - assert len({step.dst_rank for step in steps if step.src_rank == source_rank}) > 1 - if require_redundant_source: - assert any(step.src_local_row >= 32 for step in steps) - - -def _depth_worker(rank, port): - os.environ["MASTER_ADDR"] = "127.0.0.1" - os.environ["MASTER_PORT"] = str(port) - torch.cuda.set_device(rank) - dist.init_process_group("gloo", rank=rank, world_size=8) - control_group = dist.new_group(list(range(8)), backend="gloo") - transfer_group = dist.new_group(list(range(8)), backend="gloo") - initial_placement = build_initial_redundant_expert_ids(256, 8, 4) - weights = [_DepthWeight(rank, layer_index, initial_placement) for layer_index in range(9)] - transfer = NixlEPLBTransfer(weights, transfer_group, rank, world_size=8) - assert transfer.staging_depth == 8 - staging_bytes = sum(tensor.nbytes for staging in transfer.staging for _, tensor in staging) - one_layer_staging_bytes = sum(tensor[:4].nbytes for _, tensor in transfer.live[0]) - assert staging_bytes == 8 * one_layer_staging_bytes - current = initial_placement - first_order = [8, 0, 5, 1, 7, 3, 6, 2, 4] - first_targets = { - layer_index: align_target_placement(current, _depth_target(layer_index)) for layer_index in range(9) - } - first_plans = [ - (layer_index, build_transfer_plan(current, first_targets[layer_index], 256, 8, 8)) - for layer_index in first_order - ] - _assert_peer_coverage(first_plans) - first_snapshot = _clone_depth_live(weights) - transfer.start(first_plans, transfer.prepare_transfer(first_plans)) - committed = 0 - pending = _wait_all_pending(transfer, control_group, 8) - _assert_depth_staging(rank, first_plans[:8], first_targets, pending, transfer) - _assert_depth_snapshot(weights, first_snapshot) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) - transfer.commit(layer_index, buffer_index) - committed += 1 - pending = _wait_all_pending(transfer, control_group, 1) - _assert_depth_staging(rank, first_plans[8:], first_targets, pending, transfer) - _assert_depth_snapshot(weights, first_snapshot, [first_plans[8][0]]) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - _assert_depth_snapshot(weights, first_snapshot, [layer for layer, _ in first_plans[committed:]]) - transfer.commit(layer_index, buffer_index) - committed += 1 - transfer.finish() - torch.cuda.synchronize() - dist.barrier(group=control_group) - _assert_depth_live(weights, rank, first_plans, first_targets) - - second_order = [7, 2, 4] - second_targets = { - layer_index: align_target_placement( - first_targets[layer_index], - torch.tensor([[((dst + layer_index + slot + 3) % 8) * 32 + slot for slot in range(4)] for dst in range(8)]), - ) - for layer_index in second_order - } - second_plans = [ - ( - layer_index, - build_transfer_plan(first_targets[layer_index], second_targets[layer_index], 256, 8, 8), - ) - for layer_index in second_order - ] - _assert_peer_coverage(second_plans, require_redundant_source=True) - second_snapshot = _clone_depth_live(weights) - transfer.start(second_plans, transfer.prepare_transfer(second_plans)) - pending = _wait_all_pending(transfer, control_group, len(second_plans)) - _assert_depth_staging(rank, second_plans, second_targets, pending, transfer) - _assert_depth_snapshot(weights, second_snapshot) - dist.barrier(group=control_group) - if rank == 0: - time.sleep(0.1) - for layer_index, buffer_index in pending: - transfer.commit(layer_index, buffer_index) - transfer.finish() - torch.cuda.synchronize() - dist.barrier(group=control_group) - _assert_depth_live(weights, rank, second_plans, second_targets) - _assert_depth_snapshot(weights, second_snapshot, set(range(9)) - set(second_order)) - _assert_depth_snapshot(weights, second_snapshot, primary_only=True) - transfer.shutdown() - dist.barrier(group=control_group) - dist.destroy_process_group() - - -@pytest.mark.skipif( - not torch.cuda.is_available() or torch.cuda.device_count() < 8, - reason="requires eight CUDA GPUs", -) -def test_eplb_transfer_eight_gpu_bounded_staging_reuse(): - mp.spawn(_depth_worker, args=(_free_port(),), nprocs=8, join=True) diff --git a/unit_tests/models/deepseek_v4/test_memory_profile.py b/unit_tests/models/deepseek_v4/test_memory_profile.py deleted file mode 100644 index bf9aaf0ead..0000000000 --- a/unit_tests/models/deepseek_v4/test_memory_profile.py +++ /dev/null @@ -1,114 +0,0 @@ -from types import SimpleNamespace - -import pytest -import torch - -from lightllm.common.eplb_utils import extract_eplb_expert_tensors -from lightllm.common.kv_cache_mem_manager.mem_manager import MemoryManager -from lightllm.models.deepseek_v4.model import ( - DeepseekV4TpPartModel, - _get_eplb_sampling_peak_nbytes, - _get_eplb_staging_nbytes, -) -from lightllm.utils import profile_max_tokens - - -def _expert(rows=4, redundant=2): - def pack(cols, scale=True): - return SimpleNamespace( - weight=torch.empty((rows, cols), dtype=torch.uint8), - weight_scale=torch.empty((rows, 2), dtype=torch.float32) if scale else None, - ) - - counter = torch.zeros((5, 4), dtype=torch.int64) - return SimpleNamespace( - w13=pack(8), - w2=pack(4, False), - expert_parallel_state=SimpleNamespace( - eplb=SimpleNamespace(num_redundant_experts_per_rank=redundant, route_counter=counter) - ), - ) - - -@pytest.mark.parametrize( - "enable,draft,redundant,staging,sampling,exclusion", - [ - (True, False, 2, 40, 1536, 40), - (True, False, 0, 0, 1536, 0), - (False, False, 2, 0, 0, 0), - (True, True, 2, 0, 0, 0), - ], -) -def test_eplb_helpers_and_model_dedupe(enable, draft, redundant, staging, sampling, exclusion): - expert = _expert(redundant=redundant) - model = DeepseekV4TpPartModel.__new__(DeepseekV4TpPartModel) - model.is_mtp_draft_model = draft - model.args = SimpleNamespace(enable_prefill_eplb=enable) - model.trans_layers_weight = [SimpleNamespace(experts_=expert), SimpleNamespace(experts_=expert)] - weights = model._get_eplb_weights() - assert weights == ([expert] if enable and not draft else []) - assert _get_eplb_staging_nbytes(weights) == staging - assert _get_eplb_sampling_peak_nbytes(weights) == sampling - assert model.get_mtp_profile_weight_exclusion() == exclusion - assert sum(tensor[0].numel() * tensor.element_size() for _, tensor in extract_eplb_expert_tensors(expert)) == 20 - assert model._get_post_profile_memory_reservations() == { - name: value for name, value in (("eplb_staging", staging), ("eplb_sampling", sampling)) if value - } - - -@pytest.mark.parametrize("exclusion,expected", [(0, 1000), (200, 800), (None, 1000)]) -def test_mtp_profile_exclusion_adjustment(monkeypatch, exclusion, expected): - seen = [] - values = iter((100, 1100)) - monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) - monkeypatch.setattr(profile_max_tokens, "get_mtp_weight_layer_num", lambda: 1) - monkeypatch.setattr( - profile_max_tokens, "get_mtp_adjusted_mem_fraction", lambda **kw: seen.append(kw["target_weight_bytes"]) or 0.5 - ) - attrs = dict( - max_total_token_num=None, - is_mtp_draft_model=False, - args=SimpleNamespace(mtp_mode="x"), - config={"n_layer": 1}, - mem_fraction=0.8, - ) - if exclusion is not None: - attrs["get_mtp_profile_weight_exclusion"] = lambda: exclusion - model = SimpleNamespace(**attrs) - with profile_max_tokens.profile_mtp_weight_memory(model): - pass - assert seen == [expected] - - -@pytest.mark.parametrize("exclusion", [-1, 1001]) -def test_mtp_profile_exclusion_validation(monkeypatch, exclusion): - values = iter((100, 1100)) - monkeypatch.setattr(profile_max_tokens.torch.cuda, "memory_allocated", lambda: next(values)) - model = SimpleNamespace( - max_total_token_num=None, - is_mtp_draft_model=False, - args=SimpleNamespace(mtp_mode="x"), - config={"n_layer": 1}, - mem_fraction=0.8, - get_mtp_profile_weight_exclusion=lambda: exclusion, - ) - with pytest.raises(ValueError, match="invalid MTP profile exclusion"): - with profile_max_tokens.profile_mtp_weight_memory(model): - pass - - -@pytest.mark.parametrize("reservations,expected", [({}, 252), ({"x": 20}, 247)]) -def test_memory_manager_profile_reservation_once(monkeypatch, reservations, expected): - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.torch.cuda.empty_cache", lambda: None) - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.dist.get_world_size", lambda: 1) - monkeypatch.setattr( - "lightllm.common.kv_cache_mem_manager.mem_manager.get_available_gpu_memory", lambda w: 1024 / 1024 ** 3 - ) - monkeypatch.setattr("lightllm.common.kv_cache_mem_manager.mem_manager.get_total_gpu_memory", lambda: 0) - m = MemoryManager.__new__(MemoryManager) - m.size = None - m.memory_reservations = reservations - m.get_cell_size = lambda: 4 - m.get_fixed_memory_size = lambda: 16 - m.profile_size(1) - assert m.size == expected diff --git a/unit_tests/models/deepseek_v4/test_model.py b/unit_tests/models/deepseek_v4/test_model.py new file mode 100644 index 0000000000..f14cfc956c --- /dev/null +++ b/unit_tests/models/deepseek_v4/test_model.py @@ -0,0 +1,57 @@ +from types import SimpleNamespace + +import pytest +import torch + +from lightllm.models.deepseek_v4.model import DeepseekV4TpPartModel + + +def test_memory_check_rejects_sequence_larger_than_full_token_pool(): + model = DeepseekV4TpPartModel.__new__(DeepseekV4TpPartModel) + model.mem_manager = SimpleNamespace(size=60416) + model.batch_max_tokens = 8192 + model.max_seq_length = 262169 + model.args = SimpleNamespace(performance_mode=None) + + with pytest.raises(AssertionError, match="max_total_token_num must be >= max_seq_length"): + model._check_mem_size() + + +def test_memory_check_still_requires_one_prefill_batch(): + model = DeepseekV4TpPartModel.__new__(DeepseekV4TpPartModel) + model.mem_manager = SimpleNamespace(size=8192) + model.batch_max_tokens = 8192 + model.max_seq_length = 262169 + model.args = SimpleNamespace(performance_mode=None) + + with pytest.raises(AssertionError, match="greater than batch_max_tokens"): + model._check_mem_size() + + +def test_decode_autotune_uses_current_mtp_capacity_api(monkeypatch): + lengths = [] + calls = [] + monkeypatch.setattr("lightllm.common.basemodel.basemodel.Autotuner.start_autotune_warmup", lambda *_: None) + monkeypatch.setattr("lightllm.common.basemodel.basemodel.Autotuner.end_autotune_warmup", lambda: None) + monkeypatch.setattr(torch.distributed, "barrier", lambda: None) + monkeypatch.setattr(torch.cuda, "empty_cache", lambda: None) + monkeypatch.setattr( + "lightllm.common.basemodel.basemodel.tqdm", + lambda values, **kwargs: lengths.extend(values) or [], + ) + + model = DeepseekV4TpPartModel.__new__(DeepseekV4TpPartModel) + model.args = SimpleNamespace(run_mode="decode") + model.batch_max_tokens = 8192 + model.max_req_num = 4 + model.is_mtp_draft_model = False + model.mtp_manager = SimpleNamespace( + get_decode_tokens_per_request=lambda is_draft: calls.append(is_draft) or 4, + ) + model.layers_num = 1 + model.autotune_layers = lambda: 1 + + model._autotune_warmup() + + assert calls == [False] + assert lengths == [16, 8, 4, 1] diff --git a/unit_tests/models/deepseek_v4/test_moe_topk.py b/unit_tests/models/deepseek_v4/test_moe_topk.py deleted file mode 100644 index 1236684bc8..0000000000 --- a/unit_tests/models/deepseek_v4/test_moe_topk.py +++ /dev/null @@ -1,91 +0,0 @@ -import pytest -import torch -import torch.nn.functional as F - -from lightllm.models.deepseek_v4.triton_kernel.moe_topk import deepseek_v4_eplb_topk - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -@pytest.mark.parametrize("is_hash", [False, True]) -@pytest.mark.parametrize("record_load", [False, True]) -@pytest.mark.parametrize("token_num", [1, 4]) -@pytest.mark.parametrize("return_logical_ids", [False, True]) -def test_deepseek_v4_eplb_topk_matches_reference(is_hash, record_load, token_num, return_logical_ids): - torch.manual_seed(0) - device = "cuda" - num_experts, topk = 256, 6 - logits = torch.randn((token_num, num_experts), device=device, dtype=torch.float32) - bias = None if is_hash else torch.randn((num_experts,), device=device, dtype=torch.float32) - input_tokens = torch.arange(token_num, device=device, dtype=torch.long) - hash_indices = torch.randint(num_experts, (token_num + 4, topk), device=device) if is_hash else None - logical_to_physical = torch.arange(num_experts * 2, device=device, dtype=torch.int32).view(num_experts, 2) - replica_count = torch.full((num_experts,), 2, device=device, dtype=torch.int32) - counter = torch.zeros((3, num_experts), device=device, dtype=torch.int64) - - weights, physical_ids, logical_ids = deepseek_v4_eplb_topk( - logits=logits, - bias=bias, - input_tokens=input_tokens if is_hash else None, - hash_indices_table=hash_indices, - topk=topk, - routed_scaling_factor=1.7, - logical_to_physical_map=logical_to_physical, - logical_replica_count=replica_count, - expert_counter=counter, - sample_index=1, - record_load=record_load, - return_logical_ids=return_logical_ids, - ) - scores = F.softplus(logits).sqrt() - expected_logical_ids = hash_indices[input_tokens] if is_hash else (scores + bias).topk(topk, dim=-1).indices - expected_weights = scores.gather(1, expected_logical_ids) - expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True) * 1.7 - token_indices = torch.arange(token_num, device=device, dtype=torch.int64).view(-1, 1) - if token_num == 1: - replica_indices = torch.zeros_like(expected_logical_ids) - else: - token_phase = (token_indices * 2654435769) & 0xFFFFFFFF - expert_phase = (expected_logical_ids * 2246822519) & 0xFFFFFFFF - replica_indices = ((token_phase + expert_phase) & 0xFFFFFFFF) % 2 - expected_physical_ids = logical_to_physical[expected_logical_ids, replica_indices].to(torch.long) - torch.cuda.synchronize() - - torch.testing.assert_close(weights, expected_weights) - if return_logical_ids: - torch.testing.assert_close(logical_ids, expected_logical_ids) - else: - assert logical_ids is None - torch.testing.assert_close(physical_ids, expected_physical_ids) - expected_counter = torch.zeros_like(counter) - if record_load: - expected_counter[1].scatter_add_( - 0, - expected_logical_ids.reshape(-1), - torch.ones_like(expected_logical_ids.reshape(-1), dtype=torch.int64), - ) - torch.testing.assert_close(counter, expected_counter) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -def test_deepseek_v4_eplb_topk_empty_input(): - empty_logits = torch.empty((0, 256), device="cuda", dtype=torch.float32) - map_ = torch.arange(256, device="cuda", dtype=torch.int32).view(256, 1) - counter = torch.zeros((1, 256), device="cuda", dtype=torch.int64) - weights, physical_ids, logical_ids = deepseek_v4_eplb_topk( - logits=empty_logits, - bias=torch.zeros((256,), device="cuda"), - input_tokens=None, - hash_indices_table=None, - topk=6, - routed_scaling_factor=1.0, - logical_to_physical_map=map_, - logical_replica_count=torch.ones((256,), device="cuda", dtype=torch.int32), - expert_counter=counter, - sample_index=0, - record_load=True, - return_logical_ids=True, - ) - assert weights.shape == (0, 6) and weights.dtype == torch.float32 - assert physical_ids.shape == (0, 6) and physical_ids.dtype == torch.long - assert logical_ids.shape == (0, 6) and logical_ids.dtype == torch.long - assert counter.sum().item() == 0 diff --git a/unit_tests/models/deepseek_v4/test_vision_integration.py b/unit_tests/models/deepseek_v4/test_vision_integration.py index ad6f165306..e4ef71ca18 100644 --- a/unit_tests/models/deepseek_v4/test_vision_integration.py +++ b/unit_tests/models/deepseek_v4/test_vision_integration.py @@ -387,11 +387,10 @@ def test_router_uses_bias_vl_only_for_image_tokens(is_hash): gate_tid2eid_=SimpleNamespace(weight=hash_table), gate_bias_=SimpleNamespace(weight=text_bias), gate_bias_vl_=SimpleNamespace(weight=vision_bias), - experts_=SimpleNamespace(expert_parallel_state=None), ) infer_state = SimpleNamespace(is_prefill=True, input_ids=torch.tensor([2, 100_000], device="cuda")) - weights, indices, _ = router._select_experts(logits, infer_state, layer_weight) + weights, indices = router._select_experts(logits, infer_state, layer_weight) scores = torch.sqrt(torch.nn.functional.softplus(logits)) text_indices = hash_table[infer_state.input_ids[0]] if is_hash else (scores[0] + text_bias).topk(6).indices @@ -403,25 +402,6 @@ def test_router_uses_bias_vl_only_for_image_tokens(is_hash): torch.testing.assert_close(weights, expected, rtol=2e-5, atol=1e-6) -def test_eplb_rejects_vision_routing(): - from lightllm.models.deepseek_v4.layer_infer.transformer_layer_infer import DeepseekV4TransformerLayerInfer - - router = DeepseekV4TransformerLayerInfer.__new__(DeepseekV4TransformerLayerInfer) - router.has_vision = True - router.is_hash = False - router.vocab_size = 32 - logits = torch.zeros((1, 256), dtype=torch.float32) - infer_state = SimpleNamespace(is_prefill=True, input_ids=torch.tensor([100_000])) - layer_weight = SimpleNamespace( - gate_bias_=SimpleNamespace(weight=torch.zeros(256, dtype=torch.float32)), - gate_bias_vl_=SimpleNamespace(weight=torch.zeros(256, dtype=torch.float32)), - experts_=SimpleNamespace(expert_parallel_state=SimpleNamespace(eplb=object())), - ) - - with pytest.raises(RuntimeError, match="does not support vision routing"): - router._select_experts(logits, infer_state, layer_weight) - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_build_image_visibility_scatter(): from lightllm.models.deepseek_v4.triton_kernel.build_swa_index_dsv4 import build_image_visibility diff --git a/unit_tests/models/gemma4/test_tokenizer.py b/unit_tests/models/gemma4/test_tokenizer.py deleted file mode 100644 index 229f74039d..0000000000 --- a/unit_tests/models/gemma4/test_tokenizer.py +++ /dev/null @@ -1,38 +0,0 @@ -from types import SimpleNamespace - -from lightllm.models.gemma4.tokenizer import Gemma4Tokenizer -from lightllm.server.multimodal_params import ImageItem - - -class _Tokenizer: - bos_token_id = 2 - - def __call__(self, prompt, add_special_tokens=False): - return SimpleNamespace(input_ids=prompt) - - -def test_tokenizer_sets_image_block_span(): - tokenizer = Gemma4Tokenizer( - _Tokenizer(), - { - "image_token_id": 90, - "boi_token_id": 91, - "eoi_token_id": 92, - "vision_soft_tokens_per_image": 3, - }, - ) - image = ImageItem(type="image_size", data=(1, 1)) - image.token_id = 100 - image.token_num = 3 - - result = tokenizer.encode( - [7, 90, 90, 92, 8], - SimpleNamespace(images=[image]), - ) - - assert result == [7, 91, 100, 101, 102, 92, 8] - assert image.start_idx == 2 - assert image.block_start_idx == 2 - assert image.block_end_idx == 5 - assert image.to_dict()["block_start_idx"] == 2 - assert image.to_dict()["block_end_idx"] == 5 diff --git a/unit_tests/server/httpserver/test_pd_compact_transport.py b/unit_tests/server/httpserver/test_pd_compact_transport.py index 4d91c91845..984d55a65b 100644 --- a/unit_tests/server/httpserver/test_pd_compact_transport.py +++ b/unit_tests/server/httpserver/test_pd_compact_transport.py @@ -45,7 +45,6 @@ def test_pd_transport_without_logprobs_keeps_token_id_in_compact_packet(): class _Manager: args = SimpleNamespace(run_mode="decode") - cancel_pd_request_registration = MagicMock() async def generate(self, **_kwargs): yield 123, "token", metadata, finish_status diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 438bcf6d84..26c93d81e0 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -12,16 +12,11 @@ HttpServerManagerForPDMaster, ReqStatus, ) -from lightllm.server.pd_io_struct import ObjType, PD_Client_Obj +from lightllm.server.pd_io_struct import ObjType from lightllm.utils.error_utils import PDPrefillNodeStopGenToken, ServerBusyError -class _PDManagerStub: - def cancel_pd_request_registration(self, _group_request_id): - pass - - -class _FailingManager(_PDManagerStub): +class _FailingManager: args = SimpleNamespace(run_mode="prefill") async def generate(self, **_kwargs): @@ -29,7 +24,7 @@ async def generate(self, **_kwargs): yield -class _CancelledManager(_PDManagerStub): +class _CancelledManager: args = SimpleNamespace(run_mode="prefill") async def generate(self, **_kwargs): @@ -37,14 +32,14 @@ async def generate(self, **_kwargs): yield -class _SuccessfulManager(_PDManagerStub): +class _SuccessfulManager: args = SimpleNamespace(run_mode="prefill") async def generate(self, **_kwargs): yield 123, "token", {"count_output_tokens": 1}, FinishStatus(FinishStatus.FINISHED_STOP) -class _StopPrefillManager(_PDManagerStub): +class _StopPrefillManager: args = SimpleNamespace(run_mode="prefill") async def generate(self, **_kwargs): @@ -56,7 +51,7 @@ class _FatalGenerateError(BaseException): pass -class _FatalManager(_PDManagerStub): +class _FatalManager: args = SimpleNamespace(run_mode="decode") async def generate(self, **_kwargs): @@ -64,7 +59,7 @@ async def generate(self, **_kwargs): yield -class _BusyManager(_PDManagerStub): +class _BusyManager: args = SimpleNamespace(run_mode="decode") async def generate(self, **_kwargs): @@ -250,11 +245,10 @@ async def run(): manager.args = SimpleNamespace(config_server_host=None) manager.pd_manager = MagicMock() manager.timer_log = AsyncMock() - manager.metric_client = MagicMock() manager.infos_queues = None - p_node = SimpleNamespace(send_control_message=AsyncMock()) - d_node = SimpleNamespace(send_control_message=AsyncMock()) + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) req_status = ReqStatus(123, p_node, d_node) manager.req_id_to_out_inf = {123: req_status} @@ -269,8 +263,8 @@ async def run(): assert req_status.prefill_prompt_ids_event.is_set() assert req_status.up_status_event.is_set() assert manager.req_id_to_out_inf[123] is req_status - p_node.send_control_message.assert_not_awaited() - d_node.send_control_message.assert_not_awaited() + p_node.websocket.send_bytes.assert_not_awaited() + d_node.websocket.send_bytes.assert_not_awaited() with pytest.raises( RuntimeError, @@ -291,7 +285,6 @@ async def run(): manager.args = SimpleNamespace(config_server_host=None) manager.pd_manager = MagicMock() manager.timer_log = AsyncMock() - manager.metric_client = MagicMock() manager.infos_queues = None req_status = ReqStatus(123, MagicMock(), MagicMock()) @@ -353,8 +346,8 @@ def test_pd_master_abort_removes_request_even_when_node_notifications_fail(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.req_id_to_out_inf = {} - p_node = SimpleNamespace(send_control_message=AsyncMock(side_effect=ConnectionError("p down"))) - d_node = SimpleNamespace(send_control_message=AsyncMock(side_effect=ConnectionError("d down"))) + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock(side_effect=ConnectionError("p down")))) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock(side_effect=ConnectionError("d down")))) manager.req_id_to_out_inf[123] = ReqStatus(123, p_node, d_node) await manager.abort(123) @@ -368,41 +361,12 @@ def test_pd_master_abort_uses_explicit_nodes_when_request_status_is_missing(): async def run(): manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) manager.req_id_to_out_inf = {} - p_node = SimpleNamespace(send_control_message=AsyncMock()) - d_node = SimpleNamespace(send_control_message=AsyncMock()) + p_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) + d_node = SimpleNamespace(websocket=SimpleNamespace(send_bytes=AsyncMock())) await manager.abort(123, p_node=p_node, d_node=d_node) - p_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) - d_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) - - asyncio.run(run()) - - -def test_pd_master_abort_does_not_wait_for_disconnected_node_inflight_send(): - async def run(): - send_started = asyncio.Event() - release_send = asyncio.Event() - - class _BlockingWebSocket: - async def send_bytes(self, _payload): - send_started.set() - await release_send.wait() - - p_node = PD_Client_Obj(1, "prefill:8000", "prefill", {}, websocket=_BlockingWebSocket()) - d_node = SimpleNamespace(send_control_message=AsyncMock()) - inflight_send = asyncio.create_task(p_node.send_control_message(b"request")) - await send_started.wait() - - manager = HttpServerManagerForPDMaster.__new__(HttpServerManagerForPDMaster) - manager.req_id_to_out_inf = {123: ReqStatus(123, p_node, d_node)} - p_node.websocket = None - - try: - await asyncio.wait_for(manager.abort(123), timeout=1) - d_node.send_control_message.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) - finally: - release_send.set() - await inflight_send + p_node.websocket.send_bytes.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) + d_node.websocket.send_bytes.assert_awaited_once_with(pickle.dumps((ObjType.ABORT, 123))) asyncio.run(run()) diff --git a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py index 6eff400c3d..a3c95b270c 100644 --- a/unit_tests/server/httpserver/test_pd_master_cached_tokens.py +++ b/unit_tests/server/httpserver/test_pd_master_cached_tokens.py @@ -12,7 +12,7 @@ def _make_manager(monkeypatch): monkeypatch.setattr( "lightllm.server.httpserver.manager.HttpServerManager._check_and_repair_length", - lambda self, *a, **k: None, + classmethod(lambda cls, *a, **k: asyncio.sleep(0)), ) monkeypatch.setattr(SamplingParams, "from_buffer_copy", classmethod(lambda cls, other: copy.copy(other))) mgr = object.__new__(HttpServerManagerForPDMaster) diff --git a/unit_tests/server/httpserver/test_running_request_lifecycle.py b/unit_tests/server/httpserver/test_running_request_lifecycle.py index 63a8ccb7ac..6fe231a0a6 100644 --- a/unit_tests/server/httpserver/test_running_request_lifecycle.py +++ b/unit_tests/server/httpserver/test_running_request_lifecycle.py @@ -37,7 +37,7 @@ def _make_manager(mode: NodeRole): manager._log_stage_timing = MagicMock() manager._log_req_header = AsyncMock() manager._encode = AsyncMock(return_value=[10, 11, 12]) - manager._check_and_repair_length = MagicMock() + manager._check_and_repair_length = AsyncMock(side_effect=lambda prompt_ids, _params: prompt_ids) manager._release_multimodal_resources = AsyncMock() manager.abort = AsyncMock() manager._register_running_request = AsyncMock() diff --git a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py b/unit_tests/server/router/model_infer/test_ep_balance_monitor.py deleted file mode 100644 index 2ddcaafddf..0000000000 --- a/unit_tests/server/router/model_infer/test_ep_balance_monitor.py +++ /dev/null @@ -1,546 +0,0 @@ -import threading -from array import array -from types import SimpleNamespace - -import pytest -import torch -from prometheus_client import generate_latest - -from lightllm.common.basemodel.layer_weights.meta_weights.fused_moe.ep_balance import PrefillEPBalanceCounters -from lightllm.distributed import communication_op as communication_op_module -from lightllm.server.metrics.metrics import Monitor -from lightllm.server.router.model_infer.mode_backend import ep_balance_monitor as monitor_module -from lightllm.server.router.model_infer.mode_backend.ep_balance_monitor import ( - EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS, - calculate_prefill_placement_pressure_drift, - calculate_prefill_balance_stats, - classify_prefill_placement_pressure_drift, - should_enable_ep_balance_monitor, -) - - -@pytest.fixture(autouse=True) -def _mock_non_sm100(monkeypatch): - monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: False) - - -def _stats(source_token_replication: int): - return calculate_prefill_balance_stats( - torch.tensor([[[[800, 40], [800, 20]], [[800, 20], [800, 40]]]], dtype=torch.int64), - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=source_token_replication, - ) - - -def _monitor_args(**overrides): - args = { - "enable_ep_moe": True, - "disable_ep_balance_monitor": False, - "run_mode": "normal", - "enable_prefill_cudagraph": False, - } - args.update(overrides) - return SimpleNamespace(**args) - - -def _pressure_round_stats(bucket_rank_loads): - """Build [round, layer=1, rank, route/compute] CPU samples for drift tests.""" - return torch.tensor([[[[0, load] for load in rank_loads]] for rank_loads in bucket_rank_loads], dtype=torch.int64) - - -def test_pressure_drift_is_zero_for_identical_pressure(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [2, 1]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_is_one_for_complete_hot_rank_migration(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0], [0, 2]]), bucket_rounds=1) - assert drift == 1.0 - - -def test_pressure_drift_tracks_same_rank_magnitude_change(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [3, 1]]), bucket_rounds=1) - assert drift == pytest.approx(0.2) - - -def test_pressure_drift_is_invariant_to_uniform_load_scale(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 1], [4, 2]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_is_zero_for_balanced_inputs(): - drift, _ = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[8, 8], [16, 16]]), bucket_rounds=1) - assert drift == 0.0 - - -def test_pressure_drift_previous_signature_bridges_report_boundary(): - _, signature = calculate_prefill_placement_pressure_drift(_pressure_round_stats([[2, 0]]), bucket_rounds=1) - drift, next_signature = calculate_prefill_placement_pressure_drift( - _pressure_round_stats([[0, 2]]), previous_pressure_signature=signature, bucket_rounds=1 - ) - assert drift == 1.0 - assert torch.equal(next_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) - - -def test_pressure_drift_default_bucket_handles_normal_report_window(): - round_stats = _pressure_round_stats([[4, 2]] * 100) - drift, signature = calculate_prefill_placement_pressure_drift(round_stats) - assert EP_BALANCE_PRESSURE_DRIFT_BUCKET_ROUNDS == 20 - assert drift == 0.0 - assert signature.shape == (1, 2) - - -@pytest.mark.parametrize( - ("drift", "expected"), - [ - (0.0, "stable"), - (0.0999, "stable"), - (0.10, "dynamic"), - (0.2999, "dynamic"), - (0.30, "dynamic"), - (0.75, "dynamic"), - (1.0, "dynamic"), - ], -) -def test_pressure_drift_classification_boundaries(drift, expected): - assert classify_prefill_placement_pressure_drift(drift) == expected - - -def test_monitor_log_stats_reports_pressure_drift_and_bridges_reports(monkeypatch): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.layer_routed_experts = torch.tensor([1], dtype=torch.int64) - monitor.layer_flops_per_expert_token = torch.tensor([2.0], dtype=torch.float64) - monitor.layer_topks = torch.tensor([2.0], dtype=torch.float64) - monitor.source_token_replication = 1 - monitor._previous_pressure_signature = None - metric_calls = [] - monitor.metric_client = SimpleNamespace(gauge_set=lambda name, value: metric_calls.append((name, value))) - logs = [] - monkeypatch.setattr(monitor_module.logger, "info", logs.append) - - def report_for_hot_rank(hot_rank): - round_stats = torch.zeros((100, 1, 2, 2), dtype=torch.int64) - round_stats[:, :, :, monitor_module.ROUTE_LOAD] = 100 - round_stats[:, :, hot_rank, monitor_module.COMPUTE_LOAD] = 2 - monitor._log_stats(round_stats) - - report_for_hot_rank(0) - first_signature = monitor._previous_pressure_signature.clone() - report_for_hot_rank(1) - - assert "prefill_ep_placement_pressure_drift=0.0000" in logs[0] - assert "prefill_ep_placement_pressure_state=stable" in logs[0] - assert "prefill_ep_placement_pressure_drift=0.2000" in logs[1] - assert "prefill_ep_placement_pressure_state=dynamic" in logs[1] - assert torch.equal(first_signature, torch.tensor([[1.0, 0.0]], dtype=torch.float64)) - assert torch.equal(monitor._previous_pressure_signature, torch.tensor([[0.0, 1.0]], dtype=torch.float64)) - assert metric_calls == [ - ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), - ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), - ("lightllm_prefill_ep_placement_pressure_drift", 0.0), - ("lightllm_prefill_ep_critical_overhead_gflops_per_routed_token", pytest.approx(2e-11)), - ("lightllm_prefill_ep_compute_critical_overhead_ratio", pytest.approx(1.0)), - ("lightllm_prefill_ep_placement_pressure_drift", pytest.approx(0.2)), - ] - - -def test_critical_overhead_preserves_per_layer_slowest_rank(): - stats = _stats(source_token_replication=1) - assert stats is not None - assert stats["critical_overhead_gflops_per_routed_token"] == pytest.approx(7.5e-11) - assert stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx(1 / 3) - - -def test_non_tpsp_tp_replication_scales_gflops_per_token_but_not_ratio(): - tpsp_stats = _stats(source_token_replication=1) - non_tpsp_tp8_stats = _stats(source_token_replication=8) - assert tpsp_stats is not None and non_tpsp_tp8_stats is not None - assert non_tpsp_tp8_stats["critical_overhead_gflops_per_routed_token"] == pytest.approx( - tpsp_stats["critical_overhead_gflops_per_routed_token"] * 8 - ) - assert non_tpsp_tp8_stats["prefill_ep_compute_critical_overhead_ratio"] == pytest.approx( - tpsp_stats["prefill_ep_compute_critical_overhead_ratio"] - ) - - -def test_critical_overhead_is_zero_when_ranks_are_balanced(): - stats = calculate_prefill_balance_stats( - torch.tensor([[[[100, 32], [100, 32]], [[100, 64], [100, 64]]]], dtype=torch.int64), - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=1, - ) - assert stats is not None - assert stats["critical_overhead_gflops_per_routed_token"] == 0.0 - assert stats["prefill_ep_compute_critical_overhead_ratio"] == 0.0 - - -@pytest.mark.parametrize( - "round_stats", - [ - torch.tensor([[[[1, 1], [1, 1]]]], dtype=torch.int64), - torch.tensor([[[[100, 0], [100, 0]]]], dtype=torch.int64), - ], -) -def test_critical_overhead_rejects_insufficient_or_zero_compute_samples(round_stats): - assert ( - calculate_prefill_balance_stats( - round_stats, - layer_routed_experts=torch.tensor([1]), - layer_flops_per_expert_token=torch.tensor([2.0]), - layer_topks=torch.tensor([2.0]), - source_token_replication=1, - ) - is None - ) - - -def test_cpu_counter_accumulates_multiple_prefill_dispatches(): - counters = PrefillEPBalanceCounters() - counters.accumulate(route_load=3, compute_load=128) - counters.accumulate(route_load=4, compute_load=256) - assert (counters.route_load, counters.compute_load) == (7, 384) - - -def test_monitor_reuses_manager_precreated_dedicated_gloo_group(monkeypatch): - sentinel_group = object() - impl = SimpleNamespace(ep_balance_counters=None) - weight = SimpleNamespace( - fuse_moe_impl=impl, - n_routed_experts=8, - hidden_size=16, - moe_intermediate_size=32, - num_experts_per_tok=2, - ) - model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) - - monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) - monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 0) - monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) - metric_ports = [] - monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: SimpleNamespace(metric_port=4321)) - monkeypatch.setattr(monitor_module, "MetricClient", lambda port: metric_ports.append(port) or object()) - monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) - monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) - monkeypatch.setattr( - monitor_module.threading, - "Thread", - lambda *args, **kwargs: SimpleNamespace(start=lambda: None), - ) - - monitor = monitor_module.EPBalanceMonitor(model) - - assert monitor.gloo_group is sentinel_group - assert impl.ep_balance_counters is monitor.counters[0] - assert metric_ports == [4321] - - -def test_nonzero_rank_monitor_does_not_create_metric_client(monkeypatch): - sentinel_group = object() - impl = SimpleNamespace(ep_balance_counters=None) - weight = SimpleNamespace( - fuse_moe_impl=impl, - n_routed_experts=8, - hidden_size=16, - moe_intermediate_size=32, - num_experts_per_tok=2, - ) - model = SimpleNamespace(args=SimpleNamespace(enable_tpsp_mix_mode=True), tp_world_size_=1) - metric_client_calls = [] - monkeypatch.setattr(monitor_module, "_find_fused_moe_weights", lambda _: [weight]) - monkeypatch.setattr(monitor_module, "get_global_rank", lambda: 1) - monkeypatch.setattr(monitor_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(monitor_module, "dist_group_manager", SimpleNamespace(ep_balance_monitor_group=sentinel_group)) - monkeypatch.setattr(monitor_module.dist, "new_group", lambda *args, **kwargs: pytest.fail("unexpected new_group")) - monkeypatch.setattr(monitor_module, "get_shm_port_args", lambda: pytest.fail("unexpected port lookup")) - monkeypatch.setattr( - monitor_module, - "MetricClient", - lambda port: metric_client_calls.append(port) or pytest.fail("unexpected metric client"), - ) - monkeypatch.setattr( - monitor_module.threading, - "Thread", - lambda *args, **kwargs: SimpleNamespace(start=lambda: None), - ) - - monitor = monitor_module.EPBalanceMonitor(model) - - assert monitor.metric_client is None - assert metric_client_calls == [] - - -@pytest.mark.parametrize("disable_monitor", [False, True]) -def test_group_manager_creates_monitor_gloo_group_only_when_enabled(monkeypatch, disable_monitor): - monitor_group = object() - custom_groups = [] - - class FakeCustomProcessGroup: - def init_symm_mem_reduce(self): - pass - - def init_flashinfer_reduce(self): - pass - - args = SimpleNamespace( - enable_ep_moe=True, - disable_ep_balance_monitor=disable_monitor, - run_mode="normal", - enable_prefill_cudagraph=False, - disable_symm_mem_allreduce=True, - disable_flashinfer_allreduce=True, - ) - monkeypatch.setattr(communication_op_module, "get_env_start_args", lambda: args) - monkeypatch.setattr( - communication_op_module, - "CustomProcessGroup", - lambda: custom_groups.append(FakeCustomProcessGroup()) or custom_groups[-1], - ) - monkeypatch.setattr(communication_op_module, "get_global_world_size", lambda: 2) - monkeypatch.setattr(communication_op_module, "is_sm100_gpu", lambda: False) - calls = [] - monkeypatch.setattr( - communication_op_module.dist, - "new_group", - lambda *args, **kwargs: calls.append((args, kwargs)) or monitor_group, - ) - - manager = communication_op_module.DistributeGroupManager() - manager.create_groups(group_size=2) - - assert len(manager.groups) == 2 - if disable_monitor: - assert calls == [] - assert manager.ep_balance_monitor_group is None - else: - assert calls == [((), {"ranks": [0, 1], "backend": "gloo"})] - assert manager.ep_balance_monitor_group is monitor_group - - -def test_monitor_registers_prefill_ep_gauges_with_model_label(): - monitor = Monitor( - SimpleNamespace( - metric_gateway=None, - job_name="test", - grouping_key=[], - enable_monitor_auth=False, - model_name="monitor-test-model", - max_req_total_len=128, - mtp_step=0, - ) - ) - values = { - "lightllm_prefill_ep_critical_overhead_gflops_per_routed_token": 1.25, - "lightllm_prefill_ep_compute_critical_overhead_ratio": 0.3, - "lightllm_prefill_ep_placement_pressure_drift": 0.125, - } - assert set(values).issubset(monitor.monitor_registry) - for name, value in values.items(): - monitor.gauge_set(name, value) - - exposition = generate_latest(monitor.registry).decode() - for name, value in values.items(): - assert f'{name}{{model_name="monitor-test-model"}} {value}' in exposition - - -def test_record_prefill_round_stores_cumulative_counter_deltas_in_ring_buffer(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters()] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = 0 - monitor._processed_round_count = 0 - monitor._overflowed = False - - monitor.counters[0].accumulate(route_load=3, compute_load=128) - monitor.record_prefill_round() - assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (0, 0) - monitor.counters[0].accumulate(route_load=2, compute_load=256) - monitor.record_prefill_round() - - assert torch.equal( - monitor._copy_local_rounds(0, 2), - torch.tensor([[[3, 128]], [[2, 256]]], dtype=torch.int64), - ) - assert not hasattr(monitor, "_round_lock") - - -def test_spsc_ring_copy_wraps_without_a_lock(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters()] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 - monitor._processed_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2 - monitor._overflowed = False - - for value in (11, 12, 13): - monitor.counters[0].accumulate(route_load=value, compute_load=value * 10) - monitor.record_prefill_round() - - assert torch.equal( - monitor._copy_local_rounds(monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - 2, monitor._written_round_count), - torch.tensor([[[11, 110]], [[12, 120]], [[13, 130]]], dtype=torch.int64), - ) - assert not hasattr(monitor, "_round_lock") - - -def test_spsc_ring_overflow_is_deferred_to_the_monitor_thread(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.enabled = True - monitor.counters = [PrefillEPBalanceCounters(route_load=7, compute_load=70)] - monitor._round_buffer_storage = array("q", [0]) * (monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY * 2) - monitor._round_buffer = torch.frombuffer(monitor._round_buffer_storage, dtype=torch.int64).view( - monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY, 1, 2 - ) - monitor._round_ready = threading.Event() - monitor._written_round_count = monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - monitor._processed_round_count = 0 - monitor._overflowed = False - - monitor.record_prefill_round() - - assert monitor._overflowed - assert monitor._round_ready.is_set() - assert monitor._written_round_count == monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY - assert (monitor.counters[0].route_load, monitor.counters[0].compute_load) == (7, 70) - - -def test_raise_buffer_overflow_always_reports_phase_and_ring_counts(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor._written_round_count = 23 - monitor._processed_round_count = 7 - - with pytest.raises(RuntimeError) as exc_info: - monitor._raise_buffer_overflow("before_sync") - - assert str(exc_info.value) == ( - "EP balance prefill-round buffer overflowed " - f"phase=before_sync written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY}" - ) - - -def test_raise_buffer_overflow_optionally_reports_common_round_end(): - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor._written_round_count = 23 - monitor._processed_round_count = 7 - - with pytest.raises(RuntimeError) as exc_info: - monitor._raise_buffer_overflow("common_round_lag", common_round_end=19) - - assert str(exc_info.value) == ( - "EP balance prefill-round buffer overflowed " - f"phase=common_round_lag written=23 processed=7 capacity={monitor_module.EP_BALANCE_ROUND_BUFFER_CAPACITY} " - "common_round_end=19" - ) - - -def test_gather_round_stats_only_allocates_receive_buffers_on_rank_zero(monkeypatch): - local_round_stats = torch.tensor([[[3, 128]]], dtype=torch.int64) - sentinel_group = object() - - rank_zero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - rank_zero_monitor.global_rank = 0 - rank_zero_monitor.world_size = 2 - rank_zero_monitor.gloo_group = sentinel_group - - def root_gather(input_tensor, gather_list, dst, group): - assert dst == 0 and group is sentinel_group - assert len(gather_list) == 2 - gather_list[0].copy_(input_tensor) - gather_list[1].copy_(input_tensor + 1) - - monkeypatch.setattr(monitor_module.dist, "gather", root_gather) - result = rank_zero_monitor._gather_round_stats(local_round_stats) - assert torch.equal(result, torch.tensor([[[[3, 128], [4, 129]]]], dtype=torch.int64)) - - nonzero_monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - nonzero_monitor.global_rank = 1 - nonzero_monitor.world_size = 2 - nonzero_monitor.gloo_group = sentinel_group - - def nonroot_gather(input_tensor, gather_list, dst, group): - assert input_tensor is local_round_stats - assert gather_list is None - assert dst == 0 and group is sentinel_group - - monkeypatch.setattr(monitor_module.dist, "gather", nonroot_gather) - assert nonzero_monitor._gather_round_stats(local_round_stats) is None - - -def test_find_fused_moe_weights_discovers_any_layer_member_once_and_sorts(monkeypatch): - class FakeFusedMoeWeight: - def __init__(self, layer_num, enabled=True): - self.layer_num_ = layer_num - self.enable_ep_moe = enabled - - monkeypatch.setattr(monitor_module, "FusedMoeWeight", FakeFusedMoeWeight) - first = FakeFusedMoeWeight(3) - second = FakeFusedMoeWeight(1) - disabled = FakeFusedMoeWeight(0, enabled=False) - model = SimpleNamespace( - trans_layers_weight=[ - SimpleNamespace(experts_=first, alias=first, ignored=disabled), - SimpleNamespace(any_direct_member=second), - ] - ) - - assert monitor_module._find_fused_moe_weights(model) == [second, first] - - -def test_monitor_disable_detaches_counters_from_all_impls(): - impl = SimpleNamespace(ep_balance_counters="unset") - monitor = monitor_module.EPBalanceMonitor.__new__(monitor_module.EPBalanceMonitor) - monitor.weights = [SimpleNamespace(fuse_moe_impl=impl)] - monitor.enabled = True - monitor._disable() - assert impl.ep_balance_counters is None - assert not monitor.enabled - - -def test_critical_overhead_requires_minimum_samples_for_every_layer(): - round_stats = torch.tensor([[[[200, 32], [200, 32]], [[1, 32], [1, 32]]]], dtype=torch.int64) - assert ( - calculate_prefill_balance_stats( - round_stats, - layer_routed_experts=torch.tensor([1, 1]), - layer_flops_per_expert_token=torch.tensor([2.0, 4.0]), - layer_topks=torch.tensor([2.0, 2.0]), - source_token_replication=1, - ) - is None - ) - - -def test_ep_moe_normal_and_prefill_enable_monitor_by_default(): - assert should_enable_ep_balance_monitor(_monitor_args()) - assert should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill")) - - -def test_disable_ep_balance_monitor_turns_monitor_off(): - assert not should_enable_ep_balance_monitor(_monitor_args(disable_ep_balance_monitor=True)) - - -def test_non_ep_moe_and_decode_mode_do_not_enable_monitor(): - assert not should_enable_ep_balance_monitor(_monitor_args(enable_ep_moe=False)) - assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="decode")) - - -def test_prefill_cudagraph_silently_disables_monitor(): - assert not should_enable_ep_balance_monitor(_monitor_args(run_mode="prefill", enable_prefill_cudagraph=True)) - - -def test_sm100_silently_disables_monitor(monkeypatch): - monkeypatch.setattr(monitor_module, "is_sm100_gpu", lambda: True) - assert not should_enable_ep_balance_monitor(_monitor_args()) diff --git a/unit_tests/server/test_api_start_eplb.py b/unit_tests/server/test_api_start_eplb.py deleted file mode 100644 index d566a94129..0000000000 --- a/unit_tests/server/test_api_start_eplb.py +++ /dev/null @@ -1,64 +0,0 @@ -import pytest - -from lightllm.server import api_start -from lightllm.server.core.objs.start_args_type import StartArgs - - -def test_eplb_prefill_cudagraph_is_rejected_before_starting_subprocesses(monkeypatch): - args = StartArgs( - enable_ep_moe=True, - enable_prefill_eplb=True, - enable_prefill_cudagraph=True, - disable_vision=True, - disable_audio=True, - disable_shm_warning=True, - ) - - monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) - monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) - monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) - monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) - monkeypatch.setattr( - api_start.process_manager, - "start_submodule_processes", - lambda *args, **kwargs: pytest.fail("subprocess startup must not be reached"), - ) - - with pytest.raises(AssertionError, match="--enable_prefill_eplb does not support --enable_prefill_cudagraph"): - api_start._launch_subprocesses(args) - - -def test_eplb_mtp_combination_is_not_rejected_before_starting_subprocesses(monkeypatch): - args = StartArgs( - model_dir="test-model", - enable_ep_moe=True, - enable_prefill_eplb=True, - mtp_mode="vanilla_no_att", - mtp_step=1, - eos_id=0, - data_type="float16", - disable_vision=True, - disable_audio=True, - disable_shm_warning=True, - ) - - monkeypatch.setattr(api_start, "_set_envs_and_config", lambda args: None) - monkeypatch.setattr(api_start, "auto_set_max_req_total_len", lambda args: None) - monkeypatch.setattr(api_start, "auto_set_fused_shared_experts", lambda args: None) - monkeypatch.setattr(api_start, "set_unique_server_name", lambda args: None) - monkeypatch.setattr(api_start, "get_model_type", lambda model_dir: "llama") - monkeypatch.setattr(api_start, "auto_set_response_parsers", lambda args: None) - monkeypatch.setattr(api_start, "auto_configure_allreduce_flags_from_args", lambda args: None) - monkeypatch.setattr(api_start, "validate_ports", lambda ports: None) - monkeypatch.setattr(api_start, "set_env_start_args", lambda args: None) - monkeypatch.setattr(api_start, "get_shm_port_args", lambda create=False: None) - monkeypatch.setattr(api_start, "send_and_receive_node_ip", lambda args: None) - monkeypatch.setattr(api_start, "is_sm100_gpu", lambda: False) - monkeypatch.setattr( - api_start.process_manager, - "start_submodule_processes", - lambda *args, **kwargs: (object(), None), - ) - monkeypatch.setattr(api_start.process_manager, "register_process_tree", lambda process: None) - - api_start._launch_subprocesses(args) diff --git a/unit_tests/server/test_pd_master_mode.py b/unit_tests/server/test_pd_master_mode.py index 96a111ef18..6e1d006e05 100644 --- a/unit_tests/server/test_pd_master_mode.py +++ b/unit_tests/server/test_pd_master_mode.py @@ -8,7 +8,6 @@ from lightllm.server.api_cli import make_argument_parser from lightllm.server.core.objs.start_args_type import StartArgs from lightllm.server.httpserver_for_pd_master.manager import HttpServerManagerForPDMaster, PDManager -from lightllm.utils.error_utils import ServerBusyError def test_pd_node_self_request_limit_cli_defaults_to_enabled_and_can_be_disabled(): @@ -208,22 +207,6 @@ def test_pd_manager_without_connected_nodes_is_healthy(): assert asyncio.run(manager.check_pd_nodes_health()) is True -@pytest.mark.parametrize( - ("prefill_nodes", "decode_nodes"), - [([], [object()]), ([object()], [])], -) -def test_pd_manager_returns_service_unavailable_when_a_node_role_is_empty(prefill_nodes, decode_nodes): - manager = PDManager(StartArgs()) - manager.prefill_nodes = prefill_nodes - manager.decode_nodes = decode_nodes - - with pytest.raises(ServerBusyError) as exc_info: - manager.select_p_d_node("prompt", None, None) - - assert exc_info.value.status_code == 503 - assert "PD nodes unavailable" in exc_info.value.message - - def test_prefill_registration_preserves_existing_inflight_prompt_chars(): args = StartArgs() manager = PDManager(args) diff --git a/unit_tests/utils/test_envs_utils.py b/unit_tests/utils/test_envs_utils.py index f4060e6287..0cac7caa9e 100644 --- a/unit_tests/utils/test_envs_utils.py +++ b/unit_tests/utils/test_envs_utils.py @@ -7,11 +7,11 @@ ) -def test_pd_cache_high_priority_max_age_defaults_to_180_seconds(monkeypatch): +def test_pd_cache_high_priority_max_age_defaults_to_36_seconds(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MAX_AGE_SECONDS", raising=False) get_pd_cache_high_priority_max_age_seconds.cache_clear() - assert get_pd_cache_high_priority_max_age_seconds() == 180 + assert get_pd_cache_high_priority_max_age_seconds() == 36 get_pd_cache_high_priority_max_age_seconds.cache_clear() @@ -25,29 +25,29 @@ def test_pd_cache_high_priority_max_age_reads_environment_variable(monkeypatch): get_pd_cache_high_priority_max_age_seconds.cache_clear() -def test_pd_cache_high_priority_min_prompt_tokens_defaults_to_2048(monkeypatch): +def test_pd_cache_high_priority_min_prompt_tokens_defaults_to_4096(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", raising=False) get_pd_cache_high_priority_min_prompt_tokens.cache_clear() - assert get_pd_cache_high_priority_min_prompt_tokens() == 2048 + assert get_pd_cache_high_priority_min_prompt_tokens() == 4096 get_pd_cache_high_priority_min_prompt_tokens.cache_clear() def test_pd_cache_high_priority_min_prompt_tokens_reads_environment_variable(monkeypatch): - monkeypatch.setenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", "4096") + monkeypatch.setenv("LIGHTLLM_PD_CACHE_HIGH_PRIORITY_MIN_PROMPT_TOKENS", "2048") get_pd_cache_high_priority_min_prompt_tokens.cache_clear() - assert get_pd_cache_high_priority_min_prompt_tokens() == 4096 + assert get_pd_cache_high_priority_min_prompt_tokens() == 2048 get_pd_cache_high_priority_min_prompt_tokens.cache_clear() -def test_pd_node_resource_wait_timeout_defaults_to_20_seconds(monkeypatch): +def test_pd_node_resource_wait_timeout_defaults_to_10_seconds(monkeypatch): monkeypatch.delenv("LIGHTLLM_PD_NODE_RESOURCE_WAIT_TIMEOUT_SECONDS", raising=False) get_pd_node_resource_wait_timeout_seconds.cache_clear() - assert get_pd_node_resource_wait_timeout_seconds() == 20 + assert get_pd_node_resource_wait_timeout_seconds() == 10 get_pd_node_resource_wait_timeout_seconds.cache_clear() diff --git a/unit_tests/utils/test_start_utils.py b/unit_tests/utils/test_start_utils.py index e99bada338..5401c0c358 100644 --- a/unit_tests/utils/test_start_utils.py +++ b/unit_tests/utils/test_start_utils.py @@ -49,9 +49,6 @@ def wait(self, timeout=None): def test_start_submodule_processes_returns_and_manages_psutil_processes(monkeypatch): class FakePipeReader: - def close(self): - pass - def recv(self): # 子进程应在等待初始化结果之前就进入 manager,保证此时 Ctrl-C 可以清理它们。 assert len(process_manager.processes) == 2 @@ -70,7 +67,7 @@ def start(self): def is_alive(self): return True - monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), SimpleNamespace(close=lambda: None))) + monkeypatch.setattr(start_utils.mp, "Pipe", lambda duplex: (FakePipeReader(), object())) monkeypatch.setattr(start_utils.mp, "Process", FakeMpProcess) monkeypatch.setattr( start_utils.psutil, From a2a7052d1a9ffdf765e81f7c43bf59807d484f08 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 03:51:54 +0000 Subject: [PATCH 203/214] more dist.barrier --- .../model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index 86efda9586..8be392e17c 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -298,7 +298,7 @@ def kv_trans(self, trans_tasks: List["TransTask"]): checkpoint_len=end, ) - if self.backend.is_deepseek_v4 and self.backend.args.enable_cpu_cache: + if self.backend.is_deepseek_v4: # CPU-cache restore can evict source radix pages before the scheduler all-gather fences this stream. dist.barrier(group=self.backend.node_nccl_group) From c370e9213f1632423f4c9478258bdf377b18715c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 04:32:19 +0000 Subject: [PATCH 204/214] minor --- lightllm/common/basemodel/basemodel.py | 1 - .../common/basemodel/triton_kernel/fused_moe/topk_select.py | 4 ++-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 6e8e54dfad..4d5f92b6f8 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -304,7 +304,6 @@ def _init_cudagraph(self): self.graph.warmup(self) def _init_prefill_cuda_graph(self): - # Draft models use self.run_mode="normal" even when the node is decode-only. self.prefill_graph = ( None if self.args.run_mode == "decode" or not get_env_start_args().enable_prefill_cudagraph diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py index 1c01cbd638..87cda6ce16 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/topk_select.py @@ -21,7 +21,7 @@ from lightllm.utils.sgl_utils import sgl_ops from typing import Callable, List, Optional, Tuple from lightllm.common.basemodel.triton_kernel.fused_moe.softmax_topk import softmax_topk -from lightllm.common.triton_utils.autotuner import Autotuner +from lightllm.common.triton_utils.autotuner import Autotuner, AutotuneKernelType def fused_topk( @@ -170,7 +170,7 @@ def select_experts( ######################################## warning ################################################## # here is used to match autotune feature, make topk_ids more random - if Autotuner.is_autotune_warmup(): + if Autotuner.is_kernel_autotune_warmup(AutotuneKernelType.GENERAL): rand_gen = torch.Generator(device="cuda") rand_gen.manual_seed(router_logits.shape[0]) router_logits = torch.randn(size=router_logits.shape, generator=rand_gen, dtype=torch.float32, device="cuda") From dd20c6ab006705d3946c6de0ea0b52c527be9768 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 15:38:00 +0800 Subject: [PATCH 205/214] refactor: narrow DSV4 PR to model support --- lightllm/common/basemodel/basemodel.py | 5 +- .../fused_moe/fused_moe_weight.py | 9 +- .../meta_weights/fused_moe/impl/__init__.py | 6 - .../meta_weights/fused_moe/impl/mxfp4_impl.py | 45 -- .../fused_moe/grouped_fused_moe_ep.py | 2 +- .../kv_cache_mem_manager/operator/deepseek.py | 8 +- lightllm/common/quantization/__init__.py | 9 +- lightllm/common/quantization/deepgemm.py | 143 +--- ..._fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json | 110 --- ..._fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json | 110 --- .../{topk_num=6}_NVIDIA_H100_80GB_HBM3.json | 50 -- ...t16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json | 74 -- ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 74 -- ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 74 -- ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 632 ------------------ ...torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json | 146 ---- ..._num=1,use_fp8_w8a8=true}_NVIDIA_H200.json | 110 --- ..._num=6,use_fp8_w8a8=true}_NVIDIA_H200.json | 110 --- .../{topk_num=6}_NVIDIA_H200.json | 50 -- ...orch.bfloat16,topk_num=6}_NVIDIA_H200.json | 74 -- ...=64,dtype=torch.bfloat16}_NVIDIA_H200.json | 74 -- ...M=8,dtype=torch.bfloat16}_NVIDIA_H200.json | 74 -- ...out_dtype=torch.bfloat16}_NVIDIA_H200.json | 554 --------------- ...out_dtype=torch.bfloat16}_NVIDIA_H200.json | 146 ---- lightllm/distributed/communication_op.py | 204 +++--- lightllm/models/deepseek_v4/model.py | 4 - lightllm/server/api_cli.py | 14 +- lightllm/server/api_models.py | 6 +- lightllm/server/api_openai.py | 42 +- lightllm/server/api_responses.py | 7 +- lightllm/server/api_start.py | 32 +- lightllm/server/build_prompt.py | 13 +- .../server/core/objs/py_sampling_params.py | 23 +- lightllm/server/core/objs/sampling_params.py | 40 +- lightllm/server/core/objs/start_args_type.py | 4 +- lightllm/server/function_call_parser.py | 298 ++++----- lightllm/server/httpserver/async_queue.py | 15 +- lightllm/server/httpserver/manager.py | 3 +- lightllm/server/httpserver/pd_loop.py | 59 +- .../httpserver_for_pd_master/manager.py | 148 ++-- .../pd_selector/cache_aware.py | 2 +- .../multi_level_kv_cache/disk_cache_worker.py | 5 +- lightllm/server/multimodal_params.py | 1 - lightllm/server/pd_io_struct.py | 105 +-- lightllm/server/router/manager.py | 7 +- .../model_infer/mode_backend/base_backend.py | 68 +- .../chunked_prefill/impl_for_xgrammar_mode.py | 3 +- .../mode_backend/dp_backend/control_state.py | 49 +- .../dp_backend/dp_shared_kv_trans.py | 2 +- .../mode_backend/dp_backend/impl.py | 8 +- .../mode_backend/multi_level_kv_cache.py | 4 - .../pd/decode_node_impl/decode_impl.py | 2 +- .../decode_node_impl/decode_trans_process.py | 21 +- .../mode_backend/pd/nixl_kv_transporter.py | 15 +- .../server/router/req_queue/base_queue.py | 8 +- .../router/req_queue/dp_balancer/__init__.py | 5 - .../req_queue/dp_balancer/cache_aware.py | 188 ------ lightllm/server/tokenizer.py | 10 +- .../visualserver/model_infer/model_rpc.py | 3 +- lightllm/utils/config_utils.py | 16 +- lightllm/utils/device_utils.py | 5 +- lightllm/utils/envs_utils.py | 48 +- revert.md | 20 + .../static_inference/static_benchmark.py | 94 +-- .../common/test_deepseek4_paged_cache.py | 8 +- unit_tests/server/core/objs/test_req.py | 1 - .../httpserver/test_pd_compact_transport.py | 73 -- .../httpserver/test_pd_generate_error.py | 2 +- .../router/dynamic_prompt/test_radix_cache.py | 9 +- .../mode_backend/test_dp_control.py | 77 --- .../test_dp_overlap_spec_engine.py | 2 + .../mode_backend/test_multi_level_kv_cache.py | 6 +- unit_tests/server/test_pd_start_args.py | 7 +- 73 files changed, 552 insertions(+), 3923 deletions(-) delete mode 100644 lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json delete mode 100644 lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json delete mode 100644 lightllm/server/router/req_queue/dp_balancer/cache_aware.py delete mode 100644 unit_tests/server/httpserver/test_pd_compact_transport.py delete mode 100644 unit_tests/server/router/model_infer/mode_backend/test_dp_control.py diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 4d5f92b6f8..51e4c48a3f 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -330,7 +330,7 @@ def forward(self, model_input: ModelInput): model_input.to_cuda() if model_input.is_prefill: - return self._prefill(model_input) + return self._prefill(model_input=model_input) else: return self._decode(model_input) @@ -720,6 +720,7 @@ def _token_forward(self, infer_state: InferStateInfo): post_output: PostLayerOutput = self.post_infer.token_forward( last_input_embs, infer_state=infer_state, layer_weight=self.pre_post_weight ) + hidden_collector.add_final_hidden(last_input_embs) model_output = self._create_model_output(post_output, infer_state) del post_output @@ -922,6 +923,7 @@ def _overlap_tpsp_context_forward(self, infer_state: InferStateInfo, infer_state last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) g_cache_manager.cache_env_out() + hidden_collector0.add_final_hidden(last_input_embs) hidden_collector1.add_final_hidden(last_input_embs1) model_output = self._create_model_output(post_output, infer_state) @@ -963,6 +965,7 @@ def _overlap_tpsp_token_forward(self, infer_state: InferStateInfo, infer_state1: post_output, post_output1 = self.post_infer.overlap_tpsp_token_forward( last_input_embs, last_input_embs1, infer_state, infer_state1, self.pre_post_weight ) + hidden_collector0.add_final_hidden(last_input_embs) hidden_collector1.add_final_hidden(last_input_embs1) model_output = self._create_model_output(post_output, infer_state) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 758e09cdfb..5143ff4853 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -69,7 +69,6 @@ def __init__( auto_update_redundancy_expert=self.auto_update_redundancy_expert, ) self.lock = threading.Lock() - self._moe_weight_finalized = False self._create_weight() def _init_config(self, network_config: Dict[str, Any]): @@ -342,13 +341,7 @@ def verify_load(self): e_score_correction_bias_load_ok = ( True if self.e_score_correction_bias is None else getattr(self.e_score_correction_bias, "load_ok", False) ) - load_ok = weight_load_ok and per_expert_scale_load_ok and e_score_correction_bias_load_ok - if load_ok and not self._moe_weight_finalized: - finalize = getattr(self.quant_method, "finalize_moe_weight", None) - if finalize is not None: - finalize(self) - self._moe_weight_finalized = True - return load_ok + return weight_load_ok and per_expert_scale_load_ok and e_score_correction_bias_load_ok def _create_weight(self): intermediate_size = self.split_inter_size diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py index 9b32284a1a..67bb90e4ef 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/__init__.py @@ -2,15 +2,9 @@ from .triton_impl import FuseMoeTriton from .marlin_impl import FuseMoeMarlin from .deepgemm_impl import FuseMoeDeepGEMM -from .mxfp4_impl import FuseMoeMXFP4 def select_fuse_moe_impl(quant_method: QuantizationMethod, enable_ep_moe: bool): - if quant_method.method_name == "mxfp4w4a16-b32-marlin": - if enable_ep_moe: - raise RuntimeError("mxfp4w4a16-b32-marlin does not support enable_ep_moe yet") - return FuseMoeMXFP4 - if enable_ep_moe: return FuseMoeDeepGEMM diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py deleted file mode 100644 index a7e19a9c80..0000000000 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/mxfp4_impl.py +++ /dev/null @@ -1,45 +0,0 @@ -import torch -from typing import Optional - -from lightllm.common.quantization.quantize_method import WeightPack -from .triton_impl import FuseMoeTriton - - -class FuseMoeMXFP4(FuseMoeTriton): - def create_workspace(self): - return None - - def _fused_experts( - self, - input_tensor: torch.Tensor, - w13: WeightPack, - w2: WeightPack, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - router_logits: Optional[torch.Tensor] = None, - is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, - alloc_tensor_func=torch.empty, - ): - try: - from vllm.model_executor.layers.fused_moe.activation import MoEActivation - from vllm.model_executor.layers.fused_moe.experts.marlin_moe import fused_marlin_moe - from vllm.scalar_type import scalar_types - except Exception as e: - raise RuntimeError(f"MXFP4 fused MoE requires vLLM fused kernels, error={repr(e)}") from e - - return fused_marlin_moe( - hidden_states=input_tensor.contiguous(), - w1=w13.weight, - w2=w2.weight, - bias1=None, - bias2=None, - w1_scale=w13.weight_scale, - w2_scale=w2.weight_scale, - topk_weights=topk_weights.to(torch.float32).contiguous(), - topk_ids=topk_ids.to(torch.long).contiguous(), - quant_type_id=scalar_types.float4_e2m1f.id, - global_num_experts=self.n_routed_experts, - activation=MoEActivation.SILU, - clamp_limit=clamp_limit, - ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 155797b1c5..e0b03e9618 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -58,7 +58,7 @@ def check_ep_expert_dtype(quant_method: Any): "EP MoE requires --expert_dtype to be one of ['fp8', 'fp4'], " f"but the resolved fused_moe quant method is `{expert_dtype}`. " "Please start with --expert_dtype fp8 or --expert_dtype fp4. " - "Note that --expert_dtype fp4 with EP MoE is only supported on SM100 GPUs." + "Note that --expert_dtype fp4 is only supported on SM100 GPUs." ) if expert_dtype == "fp4fp8-b32-deepgemm" and not is_sm100_gpu(): raise RuntimeError( diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 92034d671f..06b5640568 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -8,9 +8,7 @@ class Deepseek2MemOperator(NormalMemOperator): def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import ( - Deepseek2MemoryManager, - ) + from lightllm.common.kv_cache_mem_manager.deepseek2_mem_manager import Deepseek2MemoryManager mem_manager: Deepseek2MemoryManager = self.mem_manager @@ -32,9 +30,7 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: class Deepseek3_2MemOperator(Deepseek2MemOperator): def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: torch.Tensor): - from lightllm.common.kv_cache_mem_manager.deepseek3_2mem_manager import ( - Deepseek3_2MemoryManager, - ) + from lightllm.common.kv_cache_mem_manager.deepseek3_2mem_manager import Deepseek3_2MemoryManager mem_manager: Deepseek3_2MemoryManager = self.mem_manager from ...basemodel.triton_kernel.kv_copy.mla_copy_kv import destindex_copy_kv diff --git a/lightllm/common/quantization/__init__.py b/lightllm/common/quantization/__init__.py index a93fb454ce..6ce4509776 100644 --- a/lightllm/common/quantization/__init__.py +++ b/lightllm/common/quantization/__init__.py @@ -14,7 +14,6 @@ EXPERT_DTYPE_TO_QUANT_TYPE = { "fp8": "fp8w8a8-b128-deepgemm", "fp4": "fp4fp8-b32-deepgemm", - "mxfp4": "mxfp4w4a16-b32-marlin", } SUPPORTED_EXPERT_DTYPES = tuple(EXPERT_DTYPE_TO_QUANT_TYPE) @@ -25,7 +24,6 @@ def __init__(self, network_config, quant_type="none", custom_cfg_path=None, expe self.quant_type = quant_type self.start_args_expert_dtype = expert_dtype self.config_expert_dtype = network_config.get("expert_dtype", None) - self.model_type = network_config.get("model_type", None) # Parse quant_cfg first so model config only fills missing per-layer fused_moe entries; # an explicit startup argument is applied afterward and overrides every layer. self._parse_custom_cfg(custom_cfg_path) @@ -60,12 +58,6 @@ def _mapping_expert_quant_method(self): expert_dtype = self.start_args_expert_dtype or self.config_expert_dtype if expert_dtype is None: return - if ( - self.start_args_expert_dtype is None - and self.config_expert_dtype == "fp4" - and self.model_type == "deepseek_v4" - ): - expert_dtype = "mxfp4" target = self._get_expert_quant_type(expert_dtype) for layer_num in range(self.layer_num): @@ -98,6 +90,7 @@ def _mapping_quant_method(self): self.quant_type = "fp8w8a8-b128-vllm" logger.info(f"select fp8w8a8-b128 quant way: {self.quant_type}") self._mapping_expert_quant_method() + elif self.hf_quantization_method == "awq": self.quant_type = "awq" if is_awq_marlin_compatible(self.hf_quantization_config): diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 5ce122e82e..7ae8cad52f 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -5,6 +5,7 @@ from lightllm.common.quantization.registry import QUANTMETHODS from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import per_token_group_quant_fp8 from lightllm.utils.log_utils import init_logger +from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -62,7 +63,7 @@ def quantize(self, weight: torch.Tensor, output: WeightPack): from lightllm.common.basemodel.triton_kernel.quantization.fp8w8a8_block_quant_kernel import weight_quant device = output.weight.device - weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=True) + weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=is_sm100_gpu()) output.weight.copy_(weight) output.weight_scale.copy_(scale) return @@ -90,7 +91,7 @@ def apply( column_major_scales=True, scale_tma_aligned=True, alloc_func=alloc_func, - use_ue8m0_scales=True, + use_ue8m0_scales=is_sm100_gpu(), ) if out is None: @@ -199,144 +200,6 @@ def _create_weight( return mm_param, mm_param_list -@QUANTMETHODS.register(["mxfp4w4a16-b32-marlin"], platform="cuda") -class MXFP4MoEQuantizationMethod(QuantizationMethod): - def __init__(self): - super().__init__() - self.block_size = 32 - self.weight_suffix = "weight" - self.weight_zero_point_suffix = None - self.weight_scale_suffix = "scale" - self.has_weight_scale = True - self.has_weight_zero_point = False - - @property - def method_name(self): - return "mxfp4w4a16-b32-marlin" - - def quantize(self, weight: torch.Tensor, output: WeightPack): - raise NotImplementedError("mxfp4w4a16-b32-marlin only loads pre-packed MXFP4 expert weights") - - def apply( - self, - input_tensor: torch.Tensor, - weight_pack: "WeightPack", - out: Optional[torch.Tensor] = None, - workspace: Optional[torch.Tensor] = None, - use_custom_tensor_mananger: bool = True, - bias: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - raise NotImplementedError("mxfp4w4a16-b32-marlin is only implemented for fused MoE expert weights") - - def _probe_marlin_layout(self, size_n: int, size_k: int, dtype: torch.dtype, device_id: int): - """用零输入走一遍真实的 per-expert repack 路径,探出 marlin 终态布局的形状与类型。 - 只调用 finalize 同款的 vllm 函数,不复刻其内部公式,杜绝形状漂移。结果按维度缓存 - (各 MoE 层同维,全程只探两次: w13 一次、w2 一次)。""" - cache_key = (size_n, size_k, dtype) - cache = getattr(self, "_marlin_layout_cache", None) - if cache is None: - cache = self._marlin_layout_cache = {} - if cache_key in cache: - return cache[cache_key] - - import vllm._custom_ops as ops - from vllm.model_executor.layers.quantization.utils.marlin_utils import ( - get_marlin_input_dtype, - marlin_permute_scales, - ) - from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( - mxfp4_marlin_process_scales, - ) - - input_dtype = get_marlin_input_dtype() - is_a_8bit = input_dtype is not None and input_dtype.itemsize == 1 - device = f"cuda:{device_id}" - qweight = torch.zeros((size_n, size_k // 2), dtype=torch.int8, device=device).view(torch.int32).T.contiguous() - marlin_qweight = ops.gptq_marlin_repack( - b_q_weight=qweight, - perm=torch.empty(0, dtype=torch.int, device=device), - size_k=size_k, - size_n=size_n, - num_bits=4, - is_a_8bit=is_a_8bit, - ) - scale = torch.zeros((size_k // self.block_size, size_n), dtype=dtype, device=device) - marlin_scale = marlin_permute_scales( - s=scale, size_k=size_k, size_n=size_n, group_size=self.block_size, is_a_8bit=is_a_8bit - ) - marlin_scale = mxfp4_marlin_process_scales(marlin_scale, input_dtype=input_dtype) - layout = ( - (tuple(marlin_qweight.shape), marlin_qweight.dtype), - (tuple(marlin_scale.shape), marlin_scale.dtype), - ) - cache[cache_key] = layout - return layout - - def _create_weight( - self, out_dims: Union[int, List[int]], in_dim: int, dtype: torch.dtype, device_id: int, num_experts: int = 1 - ) -> Tuple[WeightPack, List[WeightPack]]: - out_dim = sum(out_dims) if isinstance(out_dims, list) else out_dims - assert in_dim % self.block_size == 0, "MXFP4 scale dimension must be divisible by block_size" - expert_prefix = (num_experts,) if num_experts > 1 else () - # CPU 暂存区: load_hf_weights 灌入原始预打包 MXFP4,finalize 时 repack 进 CUDA 终态。 - weight = torch.empty(expert_prefix + (out_dim, in_dim // 2), dtype=torch.int8, device="cpu") - weight_scale = torch.empty( - expert_prefix + (out_dim, in_dim // self.block_size), dtype=torch.float8_e8m0fnu, device="cpu" - ) - mm_param = WeightPack(weight=weight, weight_scale=weight_scale) - # CUDA 终态(marlin 布局)在构造期物化,使 mem manager 的 profile 看到真实权重占用 - # ("构造即分配、load 只灌数"的框架契约,与其它 quant 方法一致;惰性到 finalize 才 - # 进卡会让空卡 profile 把 kv 池撑到挤爆权重加载)。finalize 时 repack 结果拷入。 - (w_shape, w_dtype), (s_shape, s_dtype) = self._probe_marlin_layout(out_dim, in_dim, dtype, device_id) - mm_param.marlin_weight = torch.empty((num_experts,) + w_shape, dtype=w_dtype, device=f"cuda:{device_id}") - mm_param.marlin_weight_scale = torch.empty((num_experts,) + s_shape, dtype=s_dtype, device=f"cuda:{device_id}") - mm_param_list = self._split_weight_pack( - mm_param, - weight_out_dims=out_dims, - weight_split_dim=-2, - weight_scale_out_dims=out_dims, - weight_scale_split_dim=-2, - ) - return mm_param, mm_param_list - - def finalize_moe_weight(self, moe_weight): - try: - from vllm.model_executor.layers.quantization.utils.marlin_utils_fp4 import ( - prepare_moe_mxfp4_layer_for_marlin, - ) - except Exception as e: - raise RuntimeError(f"mxfp4w4a16-b32-marlin requires vLLM MXFP4 packing utilities, error={repr(e)}") from e - - class _MXFP4Layer: - pass - - device = torch.device("cuda", moe_weight.device_id_) - layer = _MXFP4Layer() - layer.params_dtype = moe_weight.data_type_ - w13 = moe_weight.w13.weight.view(torch.uint8).to(device=device, non_blocking=True).contiguous() - w2 = moe_weight.w2.weight.view(torch.uint8).to(device=device, non_blocking=True).contiguous() - w13_scale = moe_weight.w13.weight_scale.to(device=device, non_blocking=True).contiguous() - w2_scale = moe_weight.w2.weight_scale.to(device=device, non_blocking=True).contiguous() - ( - w13_new, - w2_new, - w13_scale_new, - w2_scale_new, - _, - _, - ) = prepare_moe_mxfp4_layer_for_marlin(layer, w13, w2, w13_scale, w2_scale, None, None) - # repack 结果拷入构造期预分配的 marlin 终态 buffer(与 AWQ marlin 路径同形态), - # CPU 暂存与 repack 临时随引用释放;shape 失配会在 copy_ 处显式报错(探针保证一致)。 - moe_weight.w13.marlin_weight.copy_(w13_new) - moe_weight.w13.marlin_weight_scale.copy_(w13_scale_new) - moe_weight.w2.marlin_weight.copy_(w2_new) - moe_weight.w2.marlin_weight_scale.copy_(w2_scale_new) - moe_weight.w13.weight = moe_weight.w13.marlin_weight - moe_weight.w13.weight_scale = moe_weight.w13.marlin_weight_scale - moe_weight.w2.weight = moe_weight.w2.marlin_weight - moe_weight.w2.weight_scale = moe_weight.w2.marlin_weight_scale - - def _deepgemm_fp8_nt(a_tuple, b_tuple, out): if HAS_DEEPGEMM: if hasattr(deep_gemm, "gemm_fp8_fp8_bf16_nt"): diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index 5204097669..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,110 +0,0 @@ -{ - "12288": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "1536": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "192": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "24576": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "384": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "48": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "49152": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "6": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "600": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "6144": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "768": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "96": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index ac4ce1ba57..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,110 +0,0 @@ -{ - "1": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 4, - "num_warps": 4 - }, - "100": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "1024": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "128": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 5, - "num_warps": 4 - }, - "16": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 4, - "num_warps": 4 - }, - "2048": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 8 - }, - "256": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "32": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 5, - "num_warps": 4 - }, - "4096": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 128, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": false, - "num_stages": 5, - "num_warps": 8 - }, - "64": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 5, - "num_warps": 4 - }, - "8": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 4, - "num_warps": 4 - }, - "8192": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index 6aa8d18c54..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_align_fused:v1/{topk_num=6}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "1": { - "BLOCK_SIZE": 256, - "num_warps": 2 - }, - "100": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "1024": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "128": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "16": { - "BLOCK_SIZE": 512, - "num_warps": 4 - }, - "2048": { - "BLOCK_SIZE": 128, - "num_warps": 4 - }, - "256": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "32": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "4096": { - "BLOCK_SIZE": 256, - "num_warps": 4 - }, - "64": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "8": { - "BLOCK_SIZE": 256, - "num_warps": 2 - }, - "8192": { - "BLOCK_SIZE": 256, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index e2da8bc968..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_DIM": 256, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - }, - "100": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "1024": { - "BLOCK_DIM": 256, - "BLOCK_M": 1, - "NUM_STAGE": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "16": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "2048": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - }, - "256": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "32": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "4096": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 4, - "num_warps": 2 - }, - "64": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "8": { - "BLOCK_DIM": 64, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 2 - }, - "8192": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index e700378de1..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_SEQ": 8, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 4, - "num_warps": 4 - }, - "100": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 2 - }, - "1024": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 4, - "num_stages": 5, - "num_warps": 1 - }, - "128": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 2 - }, - "16": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 8 - }, - "2048": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 4, - "num_stages": 3, - "num_warps": 1 - }, - "256": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 2, - "num_warps": 2 - }, - "32": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 2 - }, - "4096": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 2, - "num_stages": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 2 - }, - "8": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 8 - }, - "8192": { - "BLOCK_SEQ": 8, - "HEAD_PARALLEL_NUM": 4, - "num_stages": 4, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index 588fd4a934..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 8 - }, - "100": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 1, - "num_warps": 4 - }, - "1024": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 3, - "num_warps": 4 - }, - "128": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 1, - "num_warps": 4 - }, - "16": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 4, - "num_warps": 8 - }, - "2048": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 4, - "num_warps": 2 - }, - "256": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 1 - }, - "32": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 2, - "num_warps": 8 - }, - "4096": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 1, - "num_warps": 2 - }, - "64": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 4 - }, - "8": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 8 - }, - "8192": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 3, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index fbd605f99a..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,632 +0,0 @@ -{ - "1": { - "BLOCK_M": 64, - "BLOCK_N": 32, - "NUM_STAGES": 2, - "num_warps": 8 - }, - "100": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "1024": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1152": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "12160": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1280": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "12928": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "13184": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13312": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13568": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13824": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13952": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1408": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14080": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "14336": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "14592": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "14720": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "14848": { - "BLOCK_M": 64, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14976": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "15232": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1536": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "16": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "1664": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1792": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1920": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2048": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2176": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2304": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "24192": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2432": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "24960": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25088": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25344": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "25472": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "256": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "2560": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25728": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25856": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26112": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26240": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26752": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2688": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27008": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27136": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27520": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27776": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27904": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "28032": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2816": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "28288": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "28800": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2944": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3072": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "32": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "3200": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3328": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3456": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3584": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3712": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "384": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3840": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "3968": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4096": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4224": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "4352": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4480": { - "BLOCK_M": 32, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "46336": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "46720": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "46976": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "48896": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "49152": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "49408": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50048": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50304": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50560": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50688": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50816": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50944": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "512": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "51328": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "52608": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "53248": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "53632": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54144": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54272": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54528": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54656": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "55040": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "640": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7424": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7552": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "768": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7680": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7808": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7936": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "8064": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8192": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8320": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8448": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8576": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "896": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "8960": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json deleted file mode 100644 index 4d7e8f1183..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H100_80GB_HBM3/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H100_80GB_HBM3.json +++ /dev/null @@ -1,146 +0,0 @@ -{ - "1": { - "BLOCK_M": 128, - "BLOCK_N": 128, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "100": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1024": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "12288": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "1536": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "16": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "192": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "2048": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "24576": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "256": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "32": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "384": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4096": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "48": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "49152": { - "BLOCK_M": 32, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "6": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "600": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "6144": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "768": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "8": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 1, - "num_warps": 8 - }, - "8192": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "96": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json deleted file mode 100644 index b1aae6bfba..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=256,N=4096,expert_num=256,mul_routed_weight=true,out_dtype=torch.bfloat16,topk_num=1,use_fp8_w8a8=true}_NVIDIA_H200.json +++ /dev/null @@ -1,110 +0,0 @@ -{ - "12288": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "1536": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "192": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "24576": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "384": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "48": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "49152": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "6": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "600": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "6144": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 32, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 2, - "num_warps": 4 - }, - "768": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "96": { - "BLOCK_SIZE_K": 64, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 64, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json deleted file mode 100644 index 9ffb0efd19..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/grouped_matmul:v1/{K=4096,N=512,expert_num=256,mul_routed_weight=false,out_dtype=torch.bfloat16,topk_num=6,use_fp8_w8a8=true}_NVIDIA_H200.json +++ /dev/null @@ -1,110 +0,0 @@ -{ - "1": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "100": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "1024": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 32, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "128": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 5, - "num_warps": 4 - }, - "16": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 32, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "2048": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - }, - "256": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 3, - "num_warps": 4 - }, - "32": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 4, - "num_warps": 4 - }, - "4096": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 128, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 16, - "NEED_TRANS": false, - "num_stages": 5, - "num_warps": 8 - }, - "64": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 1, - "NEED_TRANS": true, - "num_stages": 5, - "num_warps": 4 - }, - "8": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "GROUP_SIZE_M": 16, - "NEED_TRANS": true, - "num_stages": 4, - "num_warps": 4 - }, - "8192": { - "BLOCK_SIZE_K": 128, - "BLOCK_SIZE_M": 64, - "BLOCK_SIZE_N": 128, - "GROUP_SIZE_M": 1, - "NEED_TRANS": false, - "num_stages": 3, - "num_warps": 4 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json deleted file mode 100644 index 85a20d9b1b..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_align_fused:v1/{topk_num=6}_NVIDIA_H200.json +++ /dev/null @@ -1,50 +0,0 @@ -{ - "1": { - "BLOCK_SIZE": 128, - "num_warps": 1 - }, - "100": { - "BLOCK_SIZE": 128, - "num_warps": 4 - }, - "1024": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "128": { - "BLOCK_SIZE": 128, - "num_warps": 4 - }, - "16": { - "BLOCK_SIZE": 256, - "num_warps": 8 - }, - "2048": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "256": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "32": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "4096": { - "BLOCK_SIZE": 128, - "num_warps": 8 - }, - "64": { - "BLOCK_SIZE": 128, - "num_warps": 4 - }, - "8": { - "BLOCK_SIZE": 512, - "num_warps": 4 - }, - "8192": { - "BLOCK_SIZE": 256, - "num_warps": 8 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json deleted file mode 100644 index de2f015a04..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/moe_sum_reduce:v1/{hidden_dim=4096,out_dtype=torch.bfloat16,topk_num=6}_NVIDIA_H200.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_DIM": 128, - "BLOCK_M": 1, - "NUM_STAGE": 2, - "num_warps": 4 - }, - "100": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "1024": { - "BLOCK_DIM": 256, - "BLOCK_M": 4, - "NUM_STAGE": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - }, - "16": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - }, - "2048": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 4, - "num_warps": 2 - }, - "256": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 2, - "num_warps": 1 - }, - "32": { - "BLOCK_DIM": 256, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 2 - }, - "4096": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 8 - }, - "8": { - "BLOCK_DIM": 512, - "BLOCK_M": 1, - "NUM_STAGE": 1, - "num_warps": 4 - }, - "8192": { - "BLOCK_DIM": 1024, - "BLOCK_M": 1, - "NUM_STAGE": 4, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json deleted file mode 100644 index e40e19975e..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=64,dtype=torch.bfloat16}_NVIDIA_H200.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 4, - "num_warps": 8 - }, - "100": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 3, - "num_warps": 1 - }, - "1024": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 4, - "num_stages": 1, - "num_warps": 2 - }, - "128": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 3, - "num_warps": 1 - }, - "16": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 4 - }, - "2048": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 2, - "num_stages": 5, - "num_warps": 1 - }, - "256": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 2 - }, - "32": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 5, - "num_warps": 4 - }, - "4096": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 3, - "num_warps": 1 - }, - "64": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 1, - "num_warps": 1 - }, - "8": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 2, - "num_warps": 1 - }, - "8192": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 2, - "num_warps": 2 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json deleted file mode 100644 index 88742a0b13..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/rotary_emb_fwd:v1/{HEAD_DIM=64,K_HEAD_NUM=0,Q_HEAD_NUM=8,dtype=torch.bfloat16}_NVIDIA_H200.json +++ /dev/null @@ -1,74 +0,0 @@ -{ - "1": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 16, - "num_stages": 4, - "num_warps": 1 - }, - "100": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 2 - }, - "1024": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 1, - "num_warps": 4 - }, - "128": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 2 - }, - "16": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 4 - }, - "2048": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 1, - "num_warps": 2 - }, - "256": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 2 - }, - "32": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 5, - "num_warps": 8 - }, - "4096": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 2, - "num_warps": 1 - }, - "64": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 1, - "num_warps": 4 - }, - "8": { - "BLOCK_SEQ": 1, - "HEAD_PARALLEL_NUM": 8, - "num_stages": 2, - "num_warps": 8 - }, - "8192": { - "BLOCK_SEQ": 2, - "HEAD_PARALLEL_NUM": 1, - "num_stages": 2, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json deleted file mode 100644 index d0ea86fac3..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=2048,out_dtype=torch.bfloat16}_NVIDIA_H200.json +++ /dev/null @@ -1,554 +0,0 @@ -{ - "1": { - "BLOCK_M": 256, - "BLOCK_N": 64, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "100": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "1024": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "1152": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "12160": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "128": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1280": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "13056": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13312": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13696": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13824": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "13952": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1408": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14080": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14336": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14464": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14720": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14848": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "14976": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "15104": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "1536": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "16": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "1664": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1920": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2048": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2176": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2304": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "24192": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2432": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "24960": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25088": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25344": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "256": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "2560": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25600": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25856": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "25984": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26112": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26368": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "26624": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2688": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27136": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27520": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "27776": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2816": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "28160": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "28800": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "2944": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3072": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "32": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "3200": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3328": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3456": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3584": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3712": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "384": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3840": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "3968": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4096": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4224": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4352": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "4480": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "46336": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "46592": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "48896": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "49152": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "49280": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50560": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50688": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "50944": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "512": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "51328": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "52608": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "53248": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "53632": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54400": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "54656": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "55040": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "64": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "640": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "7296": { - "BLOCK_M": 32, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7552": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "768": { - "BLOCK_M": 8, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7808": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "7936": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "8064": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8192": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8320": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8448": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "8704": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "896": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - } -} \ No newline at end of file diff --git a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json b/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json deleted file mode 100644 index fbd3649737..0000000000 --- a/lightllm/common/triton_utils/autotune_kernel_configs/triton_3.6.0/NVIDIA_H200/silu_and_mul_fwd:v1/{N=256,out_dtype=torch.bfloat16}_NVIDIA_H200.json +++ /dev/null @@ -1,146 +0,0 @@ -{ - "1": { - "BLOCK_M": 128, - "BLOCK_N": 32, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "100": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "1024": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "12288": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "128": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "1536": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "16": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "192": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "2048": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "24576": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "256": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "32": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "384": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 2, - "num_warps": 1 - }, - "4096": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "48": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 2, - "num_warps": 4 - }, - "49152": { - "BLOCK_M": 32, - "BLOCK_N": 256, - "NUM_STAGES": 4, - "num_warps": 1 - }, - "6": { - "BLOCK_M": 1, - "BLOCK_N": 32, - "NUM_STAGES": 4, - "num_warps": 8 - }, - "600": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "6144": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "64": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 4 - }, - "768": { - "BLOCK_M": 1, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "8": { - "BLOCK_M": 1, - "BLOCK_N": 64, - "NUM_STAGES": 1, - "num_warps": 1 - }, - "8192": { - "BLOCK_M": 8, - "BLOCK_N": 256, - "NUM_STAGES": 1, - "num_warps": 4 - }, - "96": { - "BLOCK_M": 1, - "BLOCK_N": 128, - "NUM_STAGES": 4, - "num_warps": 8 - } -} \ No newline at end of file diff --git a/lightllm/distributed/communication_op.py b/lightllm/distributed/communication_op.py index c6b14e990b..11573e3daf 100644 --- a/lightllm/distributed/communication_op.py +++ b/lightllm/distributed/communication_op.py @@ -20,6 +20,7 @@ import os import torch +import triton import torch.distributed as dist from torch.distributed import ReduceOp, ProcessGroup from typing import List, Dict, Optional, Set, Union @@ -43,56 +44,6 @@ logger = init_logger(__name__) -def get_deep_ep_prefill_moe_workspace_size( - num_max_tokens_per_rank: int, - hidden_size: int, - intermediate_size: int, - num_experts_per_tok: int, - num_experts: int, - world_size: int, - hidden_dtype: torch.dtype, -) -> int: - """Size one Prefill microbatch workspace for the bounded grouped-MoE path. - - The target chunk covers balanced expert routing in one pass. More skewed - routing remains correct because the consumer already splits expanded rows. - """ - tensor_alignment = 256 - expert_alignment = 128 - metadata_row_granularity = 1024 - - assert num_experts % world_size == 0 - assert intermediate_size % expert_alignment == 0 - - def align_up(value: int, alignment: int) -> int: - return (value + alignment - 1) // alignment * alignment - - hidden_bytes = torch.empty((), dtype=hidden_dtype).element_size() - max_gather_rows = align_up(world_size * num_max_tokens_per_rank, metadata_row_granularity) - num_local_experts = num_experts // world_size - target_chunk_rows = align_up( - num_max_tokens_per_rank * num_experts_per_tok + num_local_experts * (expert_alignment - 1), - expert_alignment, - ) - - gather_out = align_up(max_gather_rows * hidden_size * hidden_bytes, tensor_alignment) - silu_out = align_up(target_chunk_rows * intermediate_size * hidden_bytes, tensor_alignment) - gemm_out_a = align_up(target_chunk_rows * 2 * intermediate_size * hidden_bytes, tensor_alignment) - quant_out = align_up(target_chunk_rows * intermediate_size, tensor_alignment) - quant_scale = align_up(target_chunk_rows * (intermediate_size // expert_alignment) * 4, tensor_alignment) - gemm_out_b = align_up(target_chunk_rows * hidden_size * hidden_bytes, tensor_alignment) - - # TensorBufferManager uses first-fit allocation. W1 keeps silu_out and gemm_out_a - # together; after gemm_out_a is reused by quantization, the freed silu_out block - # can only hold gemm_out_b when it is large enough. - w1_peak = gather_out + silu_out + gemm_out_a - w2_peak = gather_out + silu_out + quant_out + quant_scale - if gemm_out_b > silu_out: - w2_peak += gemm_out_b - # TensorBufferManager may trim an unaligned prefix from the supplied view. - return max(w1_peak, w2_peak) + tensor_alignment - - try: import deep_ep @@ -157,10 +108,8 @@ def all_gather_into_tensor(self, output_: torch.Tensor, input_: torch.Tensor, as class DistributeGroupManager: def __init__(self): self.groups = [] - self.dp_control_group = None self.ep_buffer = None self.ep_low_latency_buffer = None - self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None @@ -176,8 +125,6 @@ def create_groups(self, group_size: int): if not args.disable_flashinfer_allreduce: group.init_flashinfer_reduce() self.groups.append(group) - if args.dp > 1: - self.dp_control_group = dist.new_group(ranks=list(range(get_global_world_size())), backend="gloo") return def get_default_group(self) -> CustomProcessGroup: @@ -225,16 +172,12 @@ def new_deepep_group( DeepEP legacy low-latency 路径。这里只为实际存在的执行路径分配 buffer, 避免为未使用的路径长期占用显存。 """ - args = get_env_start_args() - enable_ep_moe = args.enable_ep_moe + enable_ep_moe = get_env_start_args().enable_ep_moe prefill_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_prefill() - decode_num_max_dispatch_tokens_per_rank = ( - None if args.run_mode == "prefill" else get_deepep_num_max_dispatch_tokens_per_rank_decode() - ) + decode_num_max_dispatch_tokens_per_rank = get_deepep_num_max_dispatch_tokens_per_rank_decode() if not enable_ep_moe: self.ep_buffer = None self.ep_low_latency_buffer = None - self.ep_prefill_moe_workspace = None self.ep_mega_moe_buffer = None self.ep_num_sms = None return @@ -263,7 +206,6 @@ def new_deepep_group( ) self.ep_mega_moe_buffer = None self.ep_low_latency_buffer = None - self.ep_prefill_moe_workspace = None if not expert_quant_method_names: raise ValueError("No valid MoE quant method was found while initializing DeepEP buffers") @@ -284,12 +226,10 @@ def new_deepep_group( method_name != mega_moe_quant_method for method_name in expert_quant_method_names ) enable_mega_moe_buffer = has_mega_moe_layer + enable_low_latency_buffer = has_legacy_moe_layer else: enable_mega_moe_buffer = False - has_legacy_moe_layer = True - - enable_low_latency_buffer = has_legacy_moe_layer and args.run_mode != "prefill" - enable_prefill_workspace = has_legacy_moe_layer and args.run_mode == "prefill" + enable_low_latency_buffer = True if enable_low_latency_buffer: # FP8 MoE 的 decode 使用 legacy low-latency buffer;prefill 阶段还会将其 @@ -297,44 +237,27 @@ def new_deepep_group( decode_size_hint = deep_ep.Buffer.get_low_latency_rdma_size_hint( self.ll_decode_num_tokens, self.ll_hidden, global_world_size, self.ll_num_experts ) - num_rdma_bytes = decode_size_hint - # normal 节点同时执行 Prefill 和 Decode,复用的 RDMA buffer 必须覆盖全部 Prefill workspace。 - if args.run_mode == "normal": - workspace_size = get_deep_ep_prefill_moe_workspace_size( - num_max_tokens_per_rank=self.ll_num_tokens, - hidden_size=self.ll_hidden, - intermediate_size=moe_intermediate_size, - num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_num_experts, - world_size=global_world_size, - hidden_dtype=get_torch_dtype(args.data_type), - ) - num_rdma_bytes = max(decode_size_hint, workspace_size * len(self.groups)) + microbatch_count = len(self.groups) + min_prefill_reuse_buffer_bytes = _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( + global_world_size=global_world_size, + prefill_tokens_per_rank=prefill_num_max_dispatch_tokens_per_rank, + hidden_size=hidden_size, + moe_intermediate_size=moe_intermediate_size, + hidden_dtype=get_torch_dtype(get_env_start_args().data_type), + microbatch_count=microbatch_count, + ) + # Decode 和 prefill 不会同时使用 local RDMA storage,容量取两条路径的较大值。 + num_rdma_bytes = max(decode_size_hint, min_prefill_reuse_buffer_bytes) + # DeepEP 返回的 decode hint 不保证能被 microbatch 均分;最终再对齐一次, + # 确保每个 workspace slice 的容量及起始位置仍保持 256-byte 对齐。 + rdma_alignment = microbatch_count * 256 + num_rdma_bytes = triton.cdiv(num_rdma_bytes, rdma_alignment) * rdma_alignment self.ep_low_latency_buffer = deep_ep.Buffer( deepep_group, num_rdma_bytes=num_rdma_bytes, low_latency_mode=True, num_qps_per_rank=(self.ll_num_experts // global_world_size), ) - self.ep_prefill_moe_workspace = self.ep_low_latency_buffer.get_local_buffer_tensor( - torch.uint8, use_rdma_buffer=True - ) - - if enable_prefill_workspace: - workspace_size = get_deep_ep_prefill_moe_workspace_size( - num_max_tokens_per_rank=self.ll_num_tokens, - hidden_size=self.ll_hidden, - intermediate_size=moe_intermediate_size, - num_experts_per_tok=num_experts_per_tok, - num_experts=self.ll_num_experts, - world_size=global_world_size, - hidden_dtype=get_torch_dtype(args.data_type), - ) - self.ep_prefill_moe_workspace = torch.empty( - workspace_size * len(self.groups), - dtype=torch.uint8, - device=torch.device("cuda", torch.cuda.current_device()), - ) if enable_mega_moe_buffer: # SM100 FP4 层通过 DeepGEMM Mega MoE 完成通信和计算,不使用 legacy @@ -353,10 +276,8 @@ def new_deepep_group( moe_intermediate_size, ) logger.info( - "Initialize DeepEP MoE buffers: low_latency=%s, prefill_workspace_bytes=%s, " - "mega_moe=%s, expert_quant_method_names=%s", + "Initialize DeepEP MoE buffers: low_latency=%s, mega_moe=%s, expert_quant_method_names=%s", enable_low_latency_buffer, - self.ep_prefill_moe_workspace.numel() if self.ep_prefill_moe_workspace is not None else 0, enable_mega_moe_buffer, sorted(expert_quant_method_names), ) @@ -380,14 +301,15 @@ def _set_num_sms_for_deep_gemm(self, deepep_sms: int): logger.warning(f"set num sms for deep_gemm failed: {e}") def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch.Tensor: - """Return a slice of the workspace used by DeepEP Prefill MoE kernels. + """Return a slice of the workspace reused by DeepEP prefill MoE kernels. - Pure Prefill nodes own a dedicated workspace and do not initialize a - low-latency Decode buffer. Decode-capable nodes reuse the idle local - RDMA storage. With one communication group, the default - ``microbatch_index=0`` receives the whole workspace. With multiple - groups, the workspace is split into ``len(self.groups)`` equal slices - and each in-flight microbatch uses the slice matching its group index. + DeepEP's low-latency RDMA buffer is idle during prefill, so its local + storage is reused as temporary workspace for the expanded MoE compute + path to reduce peak GPU memory. With one communication group, the + default ``microbatch_index=0`` receives the whole workspace. With + multiple groups, the workspace is split into ``len(self.groups)`` + equal slices and each in-flight microbatch uses the slice matching its + group index. Args: microbatch_index: Zero-based microbatch and communication-group @@ -401,8 +323,8 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. initialized. The same returned slice must not be used concurrently by overlapping calls. """ - assert self.ep_prefill_moe_workspace is not None, "DeepEP Prefill MoE workspace is not initialized" - workspace = self.ep_prefill_moe_workspace + assert self.ep_low_latency_buffer is not None, "DeepEP low-latency buffer is not initialized" + workspace = self.ep_low_latency_buffer.get_local_buffer_tensor(torch.uint8, use_rdma_buffer=True) microbatch_count = len(self.groups) assert 0 <= microbatch_index < microbatch_count workspace_size = workspace.numel() // microbatch_count @@ -410,9 +332,8 @@ def get_deep_ep_prefill_moe_workspace(self, microbatch_index: int = 0) -> torch. def clear_deepep_buffer(self): """ - Decode-capable modes reuse the low-latency RDMA buffer during Prefill, - so clean it before the next low-latency Decode. Pure Prefill owns a - dedicated workspace and has no low-latency buffer to clean. + Prefill MoE compute reuses the low-latency RDMA buffer as workspace. + Clean it before the buffer is used by low-latency decode kernels. """ if self.ep_low_latency_buffer is not None: self.ep_low_latency_buffer.clean_low_latency_buffer( @@ -505,4 +426,63 @@ def _is_single_group(group: Optional[Union[ProcessGroup, CustomProcessGroup]]) - return dist.get_world_size(group=group) == 1 +def _calculate_min_chunked_expanded_moe_reuse_buffer_bytes( + global_world_size: int, + prefill_tokens_per_rank: int, + hidden_size: int, + moe_intermediate_size: int, + hidden_dtype: torch.dtype, + microbatch_count: int, +) -> int: + """计算 ``chunked_expanded_moe_forward`` 极端情况下所需的最小复用 buffer。 + + Prefill 阶段会将空闲的 legacy DeepEP local RDMA storage 复用为 expanded + grouped GEMM 的 workspace。该函数只计算这条 prefill 路径的容量下限,不包含 + low-latency decode 自身所需的 RDMA 容量。 + + 每个 prefill workspace 由全程常驻的 dense gather 输出和分块计算的临时峰值组成。 + gather 按 ``global_world_size * prefill_tokens_per_rank`` 估算最大可能接收行数, + 并向 1024 行对齐;这与运行期 workspace 探测使用的保守分桶保持一致。临时计算 + 按最小 128 行 chunk 估算:W1 阶段同时需要 ``gemm_out_a`` 和 ``silu_out``;后续 + 量化/W2 阶段涉及 ``silu_out``、FP8 量化结果、scale 及 ``gemm_out_b``。这里有意 + 沿用保守公式,避免因生命周期估计不足而低估空间。 + + 多个 microbatch/communication group 并行时,每组独占一个 workspace slice,故 + prefill 容量乘以 ``microbatch_count``。最后按 ``microbatch_count * 256`` 对齐, + 保证均分后每个 slice 仍满足 256-byte 对齐。 + """ + + def align(value: int, alignment: int) -> int: + return triton.cdiv(value, alignment) * alignment + + def tensor_bytes(rows: int, columns: int, itemsize: int) -> int: + # 沿用TensorBufferManager的256对齐规则 + return align(rows * columns * itemsize, 256) + + chunk_rows = 128 + hidden_itemsize = hidden_dtype.itemsize + scale_cols = moe_intermediate_size // 128 + + silu_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, hidden_itemsize) + gemm_out_a_bytes = tensor_bytes(chunk_rows, 2 * moe_intermediate_size, hidden_itemsize) + quant_bytes = tensor_bytes(chunk_rows, moe_intermediate_size, 1) + scale_bytes = tensor_bytes(chunk_rows, scale_cols, torch.float32.itemsize) + gemm_out_b_bytes = tensor_bytes(chunk_rows, hidden_size, hidden_itemsize) + + w1_peak_bytes = silu_bytes + gemm_out_a_bytes + + # silu 释放后,B 较小时可复用其 first-fit 空洞;B 更大则需另行预留完整空间。 + quant_w2_peak_bytes = ( + silu_bytes + quant_bytes + scale_bytes + (gemm_out_b_bytes if gemm_out_b_bytes > silu_bytes else 0) + ) + temporary_peak_bytes = max(w1_peak_bytes, quant_w2_peak_bytes) + + # 按 _get_max_chunk_rows 的 1024 行对齐规则 + gather_bytes = tensor_bytes(align(global_world_size * prefill_tokens_per_rank, 1024), hidden_size, hidden_itemsize) + per_workspace_bytes = gather_bytes + temporary_peak_bytes + + reuse_buffer_bytes = per_workspace_bytes * microbatch_count + return align(reuse_buffer_bytes, microbatch_count * 256) + + dist_group_manager = DistributeGroupManager() diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 377b1d611a..255f649b69 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -377,10 +377,6 @@ def __init__(self, tokenizer, model_dir, model_config=None): def __getattr__(self, name): return getattr(self.tokenizer, name) - @property - def xgrammar_tokenizer(self): - return self.tokenizer - def get_added_vocab(self): if self._added_vocab is None: self._added_vocab = self.tokenizer.get_added_vocab() diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index 4aee85fe6c..e1dbf0c11c 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -288,8 +288,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--dp_balancer", type=str, default="bs_balancer", - choices=["round_robin", "bs_balancer", "cache_aware"], - help="the DP balancer type; cache_aware adds token-prefix affinity, default is bs_balancer", + choices=["round_robin", "bs_balancer"], + help="the dp balancer type, default is bs_balancer", ) parser.add_argument( "--max_req_total_len", @@ -726,11 +726,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: type=str, default=None, choices=["fp8", "fp4"], - help="""Requested dtype for MoE expert weights, fp8 or fp4. Resolves the fused_moe - quant method: fp8 -> fp8w8a8-b128-deepgemm; fp4 -> fp4fp8-b32-deepgemm (online - quantization) on SM100 GPUs, or mxfp4w4a16-b32-marlin (Marlin W4A16, TP only) on other GPUs. - Defaults to `expert_dtype` in config.json if present. Per-layer override: - --quant_cfg mix_bits with name `fused_moe`.""", + help="""Expert quantization dtype for EP MoE. Supported values are + fp8 and fp4. Note that fp4 is only supported on SM100 GPUs.""", ) parser.add_argument( "--vit_quant_type", @@ -910,8 +907,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--disk_cache_dir", type=str, default=None, - help="""Base directory used to persist disk cache data. A unique service name is appended so every server - instance uses a separate subdirectory. Defaults to a temp directory when not set.""", + help="""Directory used to persist disk cache data. Defaults to a temp directory when not set.""", ) parser.add_argument( "--enable_dp_prompt_cache_fetch", diff --git a/lightllm/server/api_models.py b/lightllm/server/api_models.py index 8a39834f4d..c1f1749203 100644 --- a/lightllm/server/api_models.py +++ b/lightllm/server/api_models.py @@ -5,7 +5,7 @@ from pydantic import BaseModel, Field, field_validator, model_validator from typing import Any, Dict, List, Optional, Union, Literal, ClassVar -from lightllm.utils.config_utils import get_generation_config_diff_dict +from lightllm.utils.config_utils import get_generation_config_dict MAX_SEED = (1 << 63) - 1 @@ -163,7 +163,7 @@ class CompletionRequest(BaseModel): def load_generation_cfg(cls, weight_dir: str): """Load default values from model generation config.""" try: - generation_cfg = get_generation_config_diff_dict(weight_dir) + generation_cfg = get_generation_config_dict(weight_dir) cls._loaded_defaults = { "do_sample": generation_cfg.get("do_sample", True), "presence_penalty": generation_cfg.get("presence_penalty", 0.0), @@ -245,7 +245,7 @@ class ChatCompletionRequest(BaseModel): def load_generation_cfg(cls, weight_dir: str): """Load default values from model generation config.""" try: - generation_cfg = get_generation_config_diff_dict(weight_dir) + generation_cfg = get_generation_config_dict(weight_dir) cls._loaded_defaults = { "do_sample": generation_cfg.get("do_sample", True), "presence_penalty": generation_cfg.get("presence_penalty", 0.0), diff --git a/lightllm/server/api_openai.py b/lightllm/server/api_openai.py index c2c2bebdf6..cac8f7c3c6 100644 --- a/lightllm/server/api_openai.py +++ b/lightllm/server/api_openai.py @@ -46,6 +46,7 @@ CompletionChoice, CompletionLogprobs, CompletionStreamResponse, + CompletionStreamChoice, FunctionResponse, ToolCall, UsageInfo, @@ -171,7 +172,9 @@ def _is_force_thinking_mode(request: ChatCompletionRequest) -> bool: return False if reasoning_parser in ["qwen3-thinking", "gpt-oss", "minimax"]: return True - if reasoning_parser in ["deepseek-v3", "deepseek-v4"]: + if reasoning_parser in ["deepseek-v3"]: + return request.chat_template_kwargs is not None and request.chat_template_kwargs.get("thinking") is True + if reasoning_parser in ["deepseek-v4"]: chat_template_kwargs = request.chat_template_kwargs or {} if "thinking" in chat_template_kwargs: return chat_template_kwargs["thinking"] is True @@ -370,9 +373,6 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req sampling_params = SamplingParams() sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict) - # Chat completions don't expose output-token logprobs. PD nodes use this - # transport-only marker to avoid forwarding unused per-token metadata. - sampling_params.return_output_logprobs = False sampling_params.verify() results_generator = g_objs.httpserver_manager.generate( @@ -417,6 +417,9 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req prompt_tokens = prompt_tokens_dict[sub_ids[0]] completion_tokens = sum(count_output_tokens_dict[sub_req_id] for sub_req_id in sub_ids) cached_tokens = prompt_cache_len_dict.get(sub_ids[0], 0) + reasoning_tokens = sum( + getattr(reasoning_parser_dict.get(sub_req_id), "reasoning_tokens", 0) for sub_req_id in sub_ids + ) for i in range(request.n): sub_req_id = sub_ids[i] @@ -483,7 +486,6 @@ async def chat_completions_impl(request: ChatCompletionRequest, raw_request: Req choices.append(choice) completion_tokens_details = None if reasoning_parser: - reasoning_tokens = sum(reasoning_parser_dict[sub_req_id].reasoning_tokens for sub_req_id in sub_ids) completion_tokens_details = CompletionTokensDetails(reasoning_tokens=reasoning_tokens) usage = UsageInfo( prompt_tokens=prompt_tokens, @@ -887,7 +889,6 @@ async def completions_impl(request: CompletionRequest, raw_request: Request) -> sampling_params = SamplingParams() sampling_params.init(tokenizer=g_objs.httpserver_manager.tokenizer, **sampling_params_dict) - sampling_params.return_output_logprobs = request.logprobs is not None sampling_params.verify() # v1/completions does not support multimodal inputs, so we use an empty MultimodalParams @@ -984,22 +985,19 @@ async def stream_results() -> AsyncGenerator[bytes, None]: prompt_str = g_objs.httpserver_manager.tokenizer.decode(prompt, skip_special_tokens=False) output_text = prompt_str + output_text - stream_resp = { - "id": str(group_request_id), - "object": "text_completion", - "created": created_time, - "model": request.model, - "choices": [ - { - "text": output_text, - "index": int(choice_index), - "logprobs": None if request.logprobs is None else {}, - "finish_reason": current_finish_reason, - } - ], - "usage": None, - } - yield f"data: {json.dumps(stream_resp, ensure_ascii=False)}\n\n" + stream_choice = CompletionStreamChoice( + index=choice_index, + text=output_text, + finish_reason=current_finish_reason, + logprobs=None if request.logprobs is None else {}, + ) + stream_resp = CompletionStreamResponse( + id=group_request_id, + created=created_time, + model=request.model, + choices=[stream_choice], + ) + yield f"data: {json.dumps(stream_resp.model_dump(), ensure_ascii=False)}\n\n" usage = UsageInfo( prompt_tokens=prompt_tokens, diff --git a/lightllm/server/api_responses.py b/lightllm/server/api_responses.py index e400e9e47d..e1b79acc78 100644 --- a/lightllm/server/api_responses.py +++ b/lightllm/server/api_responses.py @@ -164,8 +164,8 @@ def _responses_to_chat_request(body: Dict[str, Any]) -> Dict[str, Any]: raise ValueError("Only truncation='disabled' is supported") effort = (body.get("reasoning") or {}).get("effort") - if effort is not None and effort not in ("low", "medium", "high"): - raise ValueError("reasoning.effort must be one of: low, medium, high") + if effort is not None: + raise ValueError("reasoning.effort is not supported") messages: List[Dict[str, Any]] = [] if body.get("instructions"): @@ -216,9 +216,6 @@ def _responses_to_chat_request(body: Dict[str, Any]) -> Dict[str, Any]: else: chat["tool_choice"] = tool_choice - if effort is not None: - chat["reasoning_effort"] = effort - response_format = _text_format_to_response_format(body) if response_format: chat["response_format"] = response_format diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 32d9f92114..8a4d2c2ca0 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -24,6 +24,7 @@ get_model_type, auto_set_fused_shared_experts, auto_set_response_parsers, + get_running_max_req_size_per_dp, ) from lightllm.utils.dist_check_utils import auto_configure_allreduce_flags_from_args @@ -137,17 +138,6 @@ def _launch_subprocesses(args: StartArgs): f"graph_max_batch_size to 32" ) - dp_size_in_node = max(1, args.dp // args.nnodes) - args.per_dp_running_max_req_size = args.running_max_req_size // dp_size_in_node - args.graph_max_batch_size = min(args.graph_max_batch_size, args.per_dp_running_max_req_size) - logger.info( - "set per-DP running request limit: global=%d, local_dp=%d, per_dp=%d, graph_max_batch_size=%d", - args.running_max_req_size, - dp_size_in_node, - args.per_dp_running_max_req_size, - args.graph_max_batch_size, - ) - if not args.disable_shm_warning: check_recommended_shm_size(args) @@ -366,14 +356,7 @@ def _launch_subprocesses(args: StartArgs): from lightllm.utils.config_utils import get_dtype args.data_type = get_dtype(args.model_dir) - assert args.data_type in [ - "fp16", - "float16", - "bf16", - "bfloat16", - "fp32", - "float32", - ] + assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] set_unique_server_name(args) @@ -399,7 +382,7 @@ def _launch_subprocesses(args: StartArgs): ) auto_configure_allreduce_flags_from_args(args) - local_request_capacity = args.per_dp_running_max_req_size + local_request_capacity = get_running_max_req_size_per_dp(args) # Limit CUDA Graph batches to the local request capacity. if not args.disable_cudagraph and args.graph_max_batch_size > local_request_capacity: @@ -621,14 +604,7 @@ def visual_only_start(args): from lightllm.utils.config_utils import get_dtype args.data_type = get_dtype(args.model_dir) - assert args.data_type in [ - "fp16", - "float16", - "bf16", - "bfloat16", - "fp32", - "float32", - ] + assert args.data_type in ["fp16", "float16", "bf16", "bfloat16", "fp32", "float32"] args.visual_node_id = uuid.uuid4().int diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py index a28e51dd9e..0b64c4af7e 100644 --- a/lightllm/server/build_prompt.py +++ b/lightllm/server/build_prompt.py @@ -148,11 +148,12 @@ async def build_prompt(request, tools) -> str: if request.role_settings: kwargs["role_setting"] = request.role_settings + if request.reasoning_effort is not None: + kwargs["reasoning_effort"] = request.reasoning_effort + if request.chat_template_kwargs: kwargs.update(request.chat_template_kwargs) - if request.reasoning_effort is not None and "reasoning_effort" not in kwargs: - kwargs["reasoning_effort"] = request.reasoning_effort # 修复一些parser类型是默认打开thinking,但是 tokenizer有时候不知道打开了thinking。导致 # 构建的reasoning parser 和 tokenizer 的行为不对齐导致的问题。 from .api_openai import _is_force_thinking_mode @@ -169,11 +170,5 @@ async def build_prompt(request, tools) -> str: try: input_str = tokenizer.apply_chat_template(**kwargs, tokenize=False, add_generation_prompt=True, tools=tools) except Exception as e: - logger.exception( - "Failed to build prompt. request=%s tools=%s template_kwargs=%s", - json.dumps(request.model_dump(by_alias=True, exclude_none=True), ensure_ascii=False, default=str), - json.dumps(tools, ensure_ascii=False, default=str), - json.dumps(kwargs, ensure_ascii=False, default=str), - ) - raise ValueError(f"Failed to build prompt: {e}") from e + raise ValueError(f"Failed to build prompt: {e}") from None return input_str diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 82e6d82427..214bdf561d 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -4,7 +4,7 @@ """ import os from typing import List, Optional, Union, Tuple -from lightllm.utils.config_utils import get_generation_config_diff_dict +from lightllm.utils.config_utils import get_generation_config_dict from lightllm.server.req_id_generator import MAX_BEST_OF from .sampling_params import MAX_SEED @@ -111,14 +111,19 @@ def __init__( @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = get_generation_config_diff_dict(weight_dir) - cls._do_sample = generation_cfg.get("do_sample", False) - cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0) - cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0) - cls._repetition_penalty = generation_cfg.get("repetition_penalty", 1.0) - cls._temperature = generation_cfg.get("temperature", 1.0) - cls._top_p = generation_cfg.get("top_p", 1.0) - cls._top_k = generation_cfg.get("top_k", -1) + generation_cfg = get_generation_config_dict(weight_dir) + + def _cfg(key, default): + value = generation_cfg.get(key) + return value if value is not None else default + + cls._do_sample = _cfg("do_sample", False) + cls._presence_penalty = _cfg("presence_penalty", 0.0) + cls._frequency_penalty = _cfg("frequency_penalty", 0.0) + cls._repetition_penalty = _cfg("repetition_penalty", 1.0) + cls._temperature = _cfg("temperature", 1.0) + cls._top_p = _cfg("top_p", 1.0) + cls._top_k = _cfg("top_k", -1) cls._stop_sequences = generation_cfg.get("stop", None) except: pass diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index c14589f3e3..545d1876fb 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -1,7 +1,7 @@ import os import ctypes from typing import Optional, List, Tuple, Union -from lightllm.utils.config_utils import get_generation_config_diff_dict +from lightllm.utils.config_utils import get_generation_config_dict from lightllm.server.req_id_generator import MAX_BEST_OF from lightllm.utils.envs_utils import get_env_start_args from .pd_kv_trans_params import PDKVTransParamObj @@ -23,11 +23,6 @@ MAX_SEED = (1 << 63) - 1 -def get_xgrammar_tokenizer(tokenizer): - """Return a tokenizer's explicit xgrammar-compatible tokenizer, if any.""" - return getattr(tokenizer, "xgrammar_tokenizer", tokenizer) - - class StopSequence(ctypes.Structure): _pack_ = 4 _fields_ = [ @@ -151,7 +146,7 @@ def initialize(self, constraint: str, tokenizer): if self.length > 0 and tokenizer is not None and constraint != "json": import xgrammar as xgr - tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(tokenizer)) + tokenizer_info = xgr.TokenizerInfo.from_huggingface(tokenizer) xgrammar_compiler = xgr.GrammarCompiler(tokenizer_info, max_threads=8) xgrammar_compiler.compile_grammar(constraint) except Exception as e: @@ -181,7 +176,7 @@ def initialize(self, constraint: str, tokenizer): if self.length > 0 and tokenizer is not None: import xgrammar as xgr - tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(tokenizer)) + tokenizer_info = xgr.TokenizerInfo.from_huggingface(tokenizer) xgrammar_compiler = xgr.GrammarCompiler(tokenizer_info, max_threads=8) xgrammar_compiler.compile_json_schema(constraint) except Exception as e: @@ -373,15 +368,11 @@ def init(self, tokenizer, **kwargs): # Initialize guided_grammar guided_grammar = kwargs.get("guided_grammar", "") - guided_json = kwargs.get("guided_json", "") - if (guided_grammar or guided_json) and get_env_start_args().output_constraint_mode != "xgrammar": - guided_grammar = "" - guided_json = "" - self.guided_grammar = GuidedGrammar() self.guided_grammar.initialize(guided_grammar, tokenizer) # Initialize guided_json + guided_json = kwargs.get("guided_json", "") self.guided_json = GuidedJsonSchema() self.guided_json.initialize(guided_json, tokenizer) @@ -415,16 +406,19 @@ def init(self, tokenizer, **kwargs): @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = get_generation_config_diff_dict(weight_dir) - cls._do_sample = generation_cfg.get("do_sample", False) - cls._presence_penalty = generation_cfg.get("presence_penalty", 0.0) - cls._frequency_penalty = generation_cfg.get("frequency_penalty", 0.0) - cls._repetition_penalty = generation_cfg.get("repetition_penalty", 1.0) - if cls._repetition_penalty is None: - cls._repetition_penalty = 1.0 - cls._temperature = generation_cfg.get("temperature", 1.0) - cls._top_p = generation_cfg.get("top_p", 1.0) - cls._top_k = generation_cfg.get("top_k", -1) + generation_cfg = get_generation_config_dict(weight_dir) + + def _cfg(key, default): + value = generation_cfg.get(key) + return value if value is not None else default + + cls._do_sample = _cfg("do_sample", False) + cls._presence_penalty = _cfg("presence_penalty", 0.0) + cls._frequency_penalty = _cfg("frequency_penalty", 0.0) + cls._repetition_penalty = _cfg("repetition_penalty", 1.0) + cls._temperature = _cfg("temperature", 1.0) + cls._top_p = _cfg("top_p", 1.0) + cls._top_k = _cfg("top_k", -1) except: pass diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 9c5260ac59..4998a7f4c8 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -54,6 +54,7 @@ class StartArgs: "qwen", "deepseekv31", "deepseekv32", + "deepseekv4", "glm47", "kimi_k2", "qwen3_coder", @@ -84,7 +85,6 @@ class StartArgs: ) chat_template: Optional[str] = field(default=None) running_max_req_size: int = field(default=256) - per_dp_running_max_req_size: Optional[int] = field(default=None, init=False) tp: int = field(default=1) dp: int = field(default=1) nnodes: int = field(default=1) @@ -226,7 +226,7 @@ class StartArgs: multinode_httpmanager_port: int = field(default=12345) disable_shm_warning: bool = field(default=False) - dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer", "cache_aware"]}) + dp_balancer: str = field(default="bs_balancer", metadata={"choices": ["round_robin", "bs_balancer"]}) enable_fused_shared_experts: bool = field(default=False) enable_mps: bool = field(default=False) multinode_router_gloo_port: int = field(default=20001) diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py index f53cef44a1..4f6aabb321 100644 --- a/lightllm/server/function_call_parser.py +++ b/lightllm/server/function_call_parser.py @@ -1481,14 +1481,11 @@ class DeepSeekV32Detector(BaseFormatDetector): Reference: https://huggingface.co/deepseek-ai/DeepSeek-V3.2 """ - def __init__(self, block_name: str = "function_calls"): + def __init__(self): super().__init__() self.dsml_token = "|DSML|" - # DeepSeek V3.2 wraps tool calls in a `function_calls` block; V4 uses - # `tool_calls`. Only the outer block name differs — the invoke/parameter - # grammar is identical — so subclasses just override block_name. - self.bot_token = f"<{self.dsml_token}{block_name}>" - self.eot_token = f"" + self.bot_token = f"<{self.dsml_token}function_calls>" + self.eot_token = f"" self.invoke_start_prefix = f"<{self.dsml_token}invoke" self.invoke_end_token = f"" self.param_end_token = f"" @@ -1513,8 +1510,6 @@ def __init__(self, block_name: str = "function_calls"): self._last_arguments = "" self._accumulated_params: List[tuple] = [] self._in_function_calls = False # Track if we're inside a function_calls block - # Text after a closed block is held unless it is whitespace before another block. - self._after_function_calls = False def has_tool_call(self, text: str) -> bool: return self.bot_token in text @@ -1534,168 +1529,138 @@ def _dsml_params_to_json(self, params: List[tuple]) -> str: def detect_and_parse(self, text: str, tools: List[Tool]) -> StreamingParseResult: """One-time parsing for DSML format tool calls.""" - first_block_start = text.find(self.bot_token) - if first_block_start == -1: - return StreamingParseResult(normal_text=text, calls=[]) + idx = text.find(self.bot_token) + normal_text = text[:idx].strip() if idx != -1 else text + if self.bot_token not in text: + return StreamingParseResult(normal_text=normal_text, calls=[]) - normal_text = text[:first_block_start].removesuffix("\n\n") + tool_indices = self._get_tool_indices(tools) calls = [] - search_pos = first_block_start - while True: - block_start = text.find(self.bot_token, search_pos) - if block_start == -1: - break - if text[search_pos:block_start].strip(): - break - - block_body_start = block_start + len(self.bot_token) - block_end = text.find(self.eot_token, block_body_start) - if block_end == -1: - break - - invoke_matches = self.invoke_regex.findall(text[block_body_start:block_end]) - for func_name, invoke_body in invoke_matches: - param_matches = self.param_regex.findall(invoke_body) - args_json = self._dsml_params_to_json(param_matches) - match_result = { - "name": func_name, - "parameters": json.loads(args_json), - } - for item in self.parse_base_json(match_result, tools): - item.tool_index = len(calls) - calls.append(item) + invoke_matches = self.invoke_regex.findall(text) + for func_name, invoke_body in invoke_matches: + if func_name not in tool_indices: + logger.warning(f"Model attempted to call undefined function: {func_name}") + continue - search_pos = block_end + len(self.eot_token) + param_matches = self.param_regex.findall(invoke_body) + args_json = self._dsml_params_to_json(param_matches) + + calls.append( + ToolCallItem( + tool_index=tool_indices[func_name], + name=func_name, + parameters=args_json, + ) + ) return StreamingParseResult(normal_text=normal_text, calls=calls) def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> StreamingParseResult: """Streaming incremental parsing for DSML format tool calls.""" - overlap = len(self.param_end_token) - 1 - has_new_param_end = self.param_end_token in self._buffer[-overlap:] + new_text self._buffer += new_text - normal_text_parts = [] - calls: List[ToolCallItem] = [] + current_text = self._buffer - try: - while True: - current_text = self._buffer + # Check if we're inside a function_calls block or starting one + has_tool = self.has_tool_call(current_text) or self._in_function_calls - if not self._in_function_calls: - block_start = current_text.find(self.bot_token) - if block_start == -1: - if self._after_function_calls: - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + if not has_tool: + partial_len = self._ends_with_partial_token(current_text, self.bot_token) + if partial_len: + return StreamingParseResult() - partial_len = self._ends_with_partial_token(current_text, self.bot_token) - if partial_len: - normal_text_parts.append(current_text[:-partial_len]) - self._buffer = current_text[-partial_len:] - else: - normal_text_parts.append(current_text) - self._buffer = "" - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - - outside_text = current_text[:block_start] - if self._after_function_calls: - if outside_text.strip(): - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) - else: - normal_text_parts.append(outside_text.removesuffix("\n\n")) - - self._buffer = current_text[block_start + len(self.bot_token) :] - self._in_function_calls = True - self._after_function_calls = False - continue + self._buffer = "" + for e_token in [self.eot_token, self.invoke_end_token]: + if e_token in new_text: + new_text = new_text.replace(e_token, "") + return StreamingParseResult(normal_text=new_text) - self._buffer = current_text.lstrip() - current_text = self._buffer - if not current_text: - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + # Mark that we're inside a function_calls block + if self.has_tool_call(current_text): + self._in_function_calls = True - if current_text.startswith(self.eot_token): - self._buffer = current_text[len(self.eot_token) :] - self._in_function_calls = False - self._after_function_calls = True - continue + # Check if function_calls block has ended + if self.eot_token in current_text: + self._in_function_calls = False - if self.eot_token.startswith(current_text): - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + if not hasattr(self, "_tool_indices"): + self._tool_indices = self._get_tool_indices(tools) - complete_invoke_match = self.invoke_regex.match(current_text) - if complete_invoke_match: - func_name = complete_invoke_match.group(1) - invoke_body = complete_invoke_match.group(2) + calls: List[ToolCallItem] = [] - if self.current_tool_id == -1: - self.current_tool_id = 0 - self.prev_tool_call_arr = [] - self.streamed_args_for_tool = [""] - self._accumulated_params = [] + try: + # Try to find complete invoke blocks first + complete_invoke_match = self.invoke_regex.search(current_text) + if complete_invoke_match: + func_name = complete_invoke_match.group(1) + invoke_body = complete_invoke_match.group(2) - while len(self.prev_tool_call_arr) <= self.current_tool_id: - self.prev_tool_call_arr.append({}) - while len(self.streamed_args_for_tool) <= self.current_tool_id: - self.streamed_args_for_tool.append("") + if self.current_tool_id == -1: + self.current_tool_id = 0 + self.prev_tool_call_arr = [] + self.streamed_args_for_tool = [""] + self._accumulated_params = [] - param_matches = self.param_regex.findall(invoke_body) - args_json = self._dsml_params_to_json(param_matches) + while len(self.prev_tool_call_arr) <= self.current_tool_id: + self.prev_tool_call_arr.append({}) + while len(self.streamed_args_for_tool) <= self.current_tool_id: + self.streamed_args_for_tool.append("") - if not self.current_tool_name_sent: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=func_name, - parameters="", - ) - ) - self.current_tool_name_sent = True + param_matches = self.param_regex.findall(invoke_body) + args_json = self._dsml_params_to_json(param_matches) - sent = len(self.streamed_args_for_tool[self.current_tool_id]) - argument_diff = args_json[sent:] - if argument_diff: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=None, - parameters=argument_diff, - ) + if not self.current_tool_name_sent: + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + name=func_name, + parameters="", ) - self.streamed_args_for_tool[self.current_tool_id] += argument_diff - - try: - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": json.loads(args_json), - } - except json.JSONDecodeError: - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": {}, - } - - self._buffer = current_text[complete_invoke_match.end() :] - self.current_tool_id += 1 - self._last_arguments = "" - self.current_tool_name_sent = False - self._accumulated_params = [] - self.streamed_args_for_tool.append("") - continue + ) + self.current_tool_name_sent = True - if self.current_tool_name_sent and not has_new_param_end: + # Send complete arguments (or remaining diff) + sent = len(self.streamed_args_for_tool[self.current_tool_id]) + argument_diff = args_json[sent:] + if argument_diff: calls.append( ToolCallItem( tool_index=self.current_tool_id, - parameters="", + name=None, + parameters=argument_diff, ) ) - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + self.streamed_args_for_tool[self.current_tool_id] += argument_diff + + try: + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": json.loads(args_json), + } + except json.JSONDecodeError: + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": {}, + } - partial_match = self.partial_invoke_regex.match(current_text) - if not partial_match: - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + # Remove processed invoke from buffer + invoke_end_pos = current_text.find(self.invoke_end_token, complete_invoke_match.start()) + if invoke_end_pos != -1: + self._buffer = current_text[invoke_end_pos + len(self.invoke_end_token) :] + else: + self._buffer = current_text[complete_invoke_match.end() :] + self.current_tool_id += 1 + self._last_arguments = "" + self.current_tool_name_sent = False + self._accumulated_params = [] + self.streamed_args_for_tool.append("") + + return StreamingParseResult(normal_text="", calls=calls) + + # Partial invoke: name is known but parameters are still streaming + partial_match = self.partial_invoke_regex.search(current_text) + if partial_match: func_name = partial_match.group(1) partial_body = partial_match.group(2) @@ -1711,28 +1676,28 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami self.streamed_args_for_tool.append("") if not self.current_tool_name_sent: - calls.append( - ToolCallItem( - tool_index=self.current_tool_id, - name=func_name, - parameters="", + if func_name in self._tool_indices: + calls.append( + ToolCallItem( + tool_index=self.current_tool_id, + name=func_name, + parameters="", + ) ) - ) - self.current_tool_name_sent = True - self.prev_tool_call_arr[self.current_tool_id] = { - "name": func_name, - "arguments": {}, - } + self.current_tool_name_sent = True + self.prev_tool_call_arr[self.current_tool_id] = { + "name": func_name, + "arguments": {}, + } else: # Stream arguments as complete parameters are parsed param_matches = self.param_regex.findall(partial_body) if param_matches and len(param_matches) > len(self._accumulated_params): self._accumulated_params = param_matches current_args_json = self._dsml_params_to_json(param_matches) - open_args_json = current_args_json[:-1] # drop trailing '}' sent = len(self.streamed_args_for_tool[self.current_tool_id]) - argument_diff = open_args_json[sent:] + argument_diff = current_args_json[sent:] if argument_diff: calls.append( @@ -1749,11 +1714,11 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami except json.JSONDecodeError: pass - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + return StreamingParseResult(normal_text="", calls=calls) except Exception as e: logger.error(f"Error in DeepSeekV32 parse_streaming_increment: {e}") - return StreamingParseResult(normal_text="".join(normal_text_parts), calls=calls) + return StreamingParseResult(normal_text="", calls=calls) class Qwen3CoderDetector(BaseFormatDetector): @@ -2049,29 +2014,12 @@ def parse_streaming_increment(self, new_text: str, tools: List[Tool]) -> Streami class DeepSeekV4Detector(DeepSeekV32Detector): - """ - Detector for DeepSeek V4 model function call format using DSML. - - Identical grammar to V3.2 (``<|DSML|invoke name="...">`` blocks with - ``<|DSML|parameter name="k" string="true|false">v`` - tags), except the outer block is named ``tool_calls`` instead of - ``function_calls`` — matching the model's own encoding (encoding_dsv4.py: - ``tool_calls_block_name = "tool_calls"``) and system prompt. - - Format Structure: - ``` - <|DSML|tool_calls> - <|DSML|invoke name="get_weather"> - <|DSML|parameter name="location" string="true">Hangzhou - - - ``` - - Reference: https://huggingface.co/deepseek-ai/DeepSeek-V4 - """ + """DeepSeek-V4 uses the V3.2 DSML payload with a tool_calls wrapper.""" def __init__(self): - super().__init__(block_name="tool_calls") + super().__init__() + self.bot_token = f"<{self.dsml_token}tool_calls>" + self.eot_token = f"" class FunctionCallParser: diff --git a/lightllm/server/httpserver/async_queue.py b/lightllm/server/httpserver/async_queue.py index 47cfed4c88..a9f0c9068f 100644 --- a/lightllm/server/httpserver/async_queue.py +++ b/lightllm/server/httpserver/async_queue.py @@ -5,6 +5,7 @@ class AsyncQueue: def __init__(self): self.datas = [] self.event = asyncio.Event() + self.lock = asyncio.Lock() async def wait_to_ready(self): try: @@ -13,15 +14,15 @@ async def wait_to_ready(self): pass async def get_all_data(self): - self.event.clear() - ans = self.datas - self.datas = [] - return ans + async with self.lock: + self.event.clear() + ans = self.datas + self.datas = [] + return ans async def put(self, obj): - was_empty = not self.datas - self.datas.append(obj) - if was_empty: + async with self.lock: + self.datas.append(obj) self.event.set() return diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index ec8949d283..1eeafea37e 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -688,7 +688,8 @@ async def _check_and_repair_length(self, prompt_ids: List[int], sampling_params: if not prompt_ids: raise InvalidRequestError("The input prompt must not be empty.") prompt_tokens = len(prompt_ids) - # MTP overlap reserves an additional KV window in get_real_supported_max_req_total_len. + # -36 用于保留通用边界余量,MTP overlap 所需的额外 KV 窗口由 + # get_real_supported_max_req_total_len 单独扣除。 real_supported_max_req_total_len = self.get_real_supported_max_req_total_len() if prompt_tokens + sampling_params.max_new_tokens > real_supported_max_req_total_len: diff --git a/lightllm/server/httpserver/pd_loop.py b/lightllm/server/httpserver/pd_loop.py index 76d61e6583..51b7eea33b 100644 --- a/lightllm/server/httpserver/pd_loop.py +++ b/lightllm/server/httpserver/pd_loop.py @@ -11,12 +11,7 @@ import sys from typing import Dict, Optional, Union, List from websockets import ClientConnection -from lightllm.server.pd_io_struct import ( - NodeRole, - ObjType, - PD_COMPACT_TOKEN_INFO_LEN, - build_pd_compact_token_info, -) +from lightllm.server.pd_io_struct import NodeRole, ObjType from lightllm.server.httpserver.async_queue import AsyncQueue from lightllm.utils.net_utils import get_hostname_ip from lightllm.utils.log_utils import init_logger @@ -239,7 +234,6 @@ async def _pd_process_generate( pd_upload_websocket: ClientConnection, pd_event: asyncio.Event, ): - return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) try: async for sub_req_id, request_output, metadata, finish_status in manager.generate( prompt=prompt, @@ -249,17 +243,7 @@ async def _pd_process_generate( pd_upload_websocket=pd_upload_websocket, pd_event=pd_event, ): - metadata.pop("prompt_ids", None) - if not return_output_logprobs: - for key in ("logprob", "cumlogprob", "special", "logprobs"): - metadata.pop(key, None) - if metadata.get("count_output_tokens") == 1: - metadata["node_mode"] = manager.args.run_mode - if not return_output_logprobs: - compact_token_info = build_pd_compact_token_info(sub_req_id, request_output, metadata, finish_status) - if compact_token_info is not None: - await forwarding_queue.put(compact_token_info) - continue + metadata["node_mode"] = manager.args.run_mode await forwarding_queue.put((sub_req_id, request_output, metadata, finish_status)) except PDPrefillNodeStopGenToken as e: logger.info(f"pd prefill node stop gen token for group_request_id {e.group_request_id}") @@ -288,49 +272,12 @@ async def _pd_process_generate( # 转发token的task async def _up_tokens_to_pd_master(forwarding_queue: AsyncQueue, websocket: ClientConnection): - max_message_size = get_lightllm_websocket_max_message_size() - while True: handle_list = await forwarding_queue.wait_to_get_all_data() if handle_list: load_info: dict = _get_load_info() - pending_handle_lists = [] - group_start = 0 - group_is_compact = len(handle_list[0]) == PD_COMPACT_TOKEN_INFO_LEN - for index in range(1, len(handle_list)): - item_is_compact = len(handle_list[index]) == PD_COMPACT_TOKEN_INFO_LEN - if item_is_compact != group_is_compact: - pending_handle_lists.append(handle_list[group_start:index]) - group_start = index - group_is_compact = item_is_compact - pending_handle_lists.append(handle_list[group_start:]) - pending_handle_lists.reverse() - while pending_handle_lists: - token_list = pending_handle_lists.pop() - obj_type = ( - ObjType.TOKEN_PACKS_COMPACT - if len(token_list[0]) == PD_COMPACT_TOKEN_INFO_LEN - else ObjType.TOKEN_PACKS - ) - payload = pickle.dumps((obj_type, token_list, load_info)) - if len(payload) <= max_message_size: - await websocket.send(payload) - continue - - if len(token_list) == 1: - raise ValueError( - f"single PD token pack is {len(payload)} bytes, exceeding websocket limit " - f"{max_message_size}" - ) - - split_index = len(token_list) // 2 - logger.warning( - f"PD token pack is {len(payload)} bytes with {len(token_list)} items, exceeding websocket " - f"limit {max_message_size}; splitting it" - ) - # 栈后进先出,先压后半段,保持 token 的原始发送顺序。 - pending_handle_lists.extend((token_list[split_index:], token_list[:split_index])) + await websocket.send(pickle.dumps((ObjType.TOKEN_PACKS, handle_list, load_info))) async def _send_heartbeat_to_pd_master(websocket: ClientConnection): diff --git a/lightllm/server/httpserver_for_pd_master/manager.py b/lightllm/server/httpserver_for_pd_master/manager.py index f00f325ea2..96f0d7203e 100644 --- a/lightllm/server/httpserver_for_pd_master/manager.py +++ b/lightllm/server/httpserver_for_pd_master/manager.py @@ -12,13 +12,7 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy()) from typing import Union, List, Tuple, Dict, Optional from lightllm.server.core.objs import FinishStatus -from ..pd_io_struct import ( - PD_Client_Obj, - PDUpKVStatus, - ObjType, - PDDecodeNodeInfo, - unpack_pd_compact_token_info, -) +from ..pd_io_struct import PD_Client_Obj, PDUpKVStatus, ObjType, PDDecodeNodeInfo from lightllm.server.core.objs import SamplingParams, StartArgs from ..multimodal_params import MultimodalParams from ..tokenizer import get_tokenizer @@ -188,9 +182,7 @@ async def _generate( self, prompt_ids=fake_prompt_ids, sampling_params=sampling_params ) - return_output_logprobs = getattr(sampling_params, "return_output_logprobs", True) origin_sampling_params = SamplingParams.from_buffer_copy(sampling_params) - origin_sampling_params.return_output_logprobs = return_output_logprobs origin_group_request_id = self.id_gen.generate_id() # Record one user request even when it is expanded into multiple independent @@ -204,7 +196,6 @@ async def _generate( generators = [] for choice_index in range(choice_count): choice_sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) - choice_sampling_params.return_output_logprobs = return_output_logprobs choice_sampling_params.n = 1 choice_sampling_params.best_of = 1 generators.append( @@ -343,7 +334,6 @@ async def _generate_one_attempt( # PD Master 吞掉该内部分段 marker,并用剩余 token 限额在同一组 P/D 节点上继续。 while remaining_max_new_tokens > 0: sampling_params = SamplingParams.from_buffer_copy(origin_sampling_params) - sampling_params.return_output_logprobs = getattr(origin_sampling_params, "return_output_logprobs", True) block_group_request_id = self.id_gen.generate_id() sampling_params.group_request_id = block_group_request_id logger.info(f"pd log gen sub req id {block_group_request_id} for main req id {origin_request_id}") @@ -370,7 +360,7 @@ async def _generate_one_attempt( results_generator = self._wait_to_token_package( p_node, d_node, - start_time if segment_index == 0 else time.time(), + start_time, block_prompt, sampling_params, multimodal_params, @@ -584,36 +574,33 @@ async def fetch_pd_stream( ready_kv_len=decode_node_info.ready_kv_len, ) - next_disconnect_check = 0.0 while True: await req_status.wait_to_ready() req_status.raise_if_error() - now = time.monotonic() - if now >= next_disconnect_check: - next_disconnect_check = now + 1.0 - if await request.is_disconnected(): - raise ClientDisconnected( - group_request_id=group_request_id, - reason="fetch_pd_stream decode period check network disconnected", - ) - token_list = req_status.pop_all_tokens() - for sub_req_id, request_output, metadata, finish_status in token_list: - output_index = metadata.get("count_output_tokens") - # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 - if output_index == 1: - if first_token_gen is False: - first_token_gen = True - node_run_mode = metadata.pop("node_mode", None) - if node_run_mode == "prefill": - if old_max_new_tokens != 1 and finish_status.is_finished_length(): - finish_status = FinishStatus(FinishStatus.NO_FINISH) + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, + reason="fetch_pd_stream decode period check network disconnected", + ) + if await req_status.can_read(self.req_id_to_out_inf): + token_list = await req_status.pop_all_tokens() + for sub_req_id, request_output, metadata, finish_status in token_list: + output_index = metadata.get("count_output_tokens") + # 因为 pd 的 prefill 和 decode 节点都有可能上报首token,所以需要做一下过滤。 + if output_index == 1: + if first_token_gen is False: + first_token_gen = True + node_run_mode = metadata.pop("node_mode", None) + if node_run_mode == "prefill": + if old_max_new_tokens != 1 and finish_status.is_finished_length(): + finish_status = FinishStatus(FinishStatus.NO_FINISH) + metadata["prompt_cache_len"] = prompt_cache_len_from_prefill + yield sub_req_id, request_output, metadata, finish_status + else: + continue + else: metadata["prompt_cache_len"] = prompt_cache_len_from_prefill yield sub_req_id, request_output, metadata, finish_status - else: - continue - else: - metadata["prompt_cache_len"] = prompt_cache_len_from_prefill - yield sub_req_id, request_output, metadata, finish_status return @@ -637,17 +624,16 @@ async def _wait_for_prefill_token_if_needed( group_request_id=group_request_id, reason="fetch_pd_stream decode period check network disconnected", ) - token_list = req_status.pop_all_tokens() - if not token_list: + if not await req_status.can_read(self.req_id_to_out_inf): continue - new_tokens.extend(token_list) + new_tokens.extend(await req_status.pop_all_tokens()) for token in new_tokens: metadata = token[2] if metadata.get("node_mode") == "prefill": prompt_cache_len = metadata.get("prompt_cache_len", 0) - req_status.put_tokens_to_front(new_tokens) + await req_status.put_tokens_to_front(new_tokens) return prompt_cache_len async def _wait_to_token_package( @@ -675,6 +661,11 @@ async def _wait_to_token_package( 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 ): + if await request.is_disconnected(): + raise ClientDisconnected( + group_request_id=group_request_id, reason="_wait_to_token_package check network disconnected" + ) + prompt_tokens = metadata["prompt_tokens"] out_token_counter += 1 prompt_cache_len = max(prompt_cache_len, metadata.get("prompt_cache_len", 0)) @@ -779,29 +770,27 @@ async def handle_loop(self): try: for obj in objs: - if obj[0] in (ObjType.TOKEN_PACKS, ObjType.TOKEN_PACKS_COMPACT): + if obj[0] == ObjType.TOKEN_PACKS: token_list, node_load_info = obj[1], obj[2] self.pd_manager.update_node_load_info(node_load_info) - compact_pack = obj[0] == ObjType.TOKEN_PACKS_COMPACT - for token_info in token_list: - if compact_pack: - sub_req_id, text, metadata, finish_status_value = unpack_pd_compact_token_info( - token_info - ) - finish_status = FinishStatus(finish_status_value) - else: - sub_req_id, text, metadata, finish_status = token_info + for sub_req_id, text, metadata, finish_status in token_list: + finish_status: FinishStatus = finish_status group_req_id = convert_sub_id_to_group_id(sub_req_id) - req_status: ReqStatus = self.req_id_to_out_inf.get(group_req_id) - if req_status is not None: - req_status.append_token((sub_req_id, text, metadata, finish_status)) + try: + req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] + async with req_status.lock: + req_status.out_token_info_list.append((sub_req_id, text, metadata, finish_status)) + req_status.event.set() + except: + pass elif obj[0] == ObjType.PD_UPLOAD_PREFILL_PROMPT_IDS: _, group_req_id, prompt_ids = obj try: req_status: ReqStatus = self.req_id_to_out_inf[group_req_id] - req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids - req_status.prefill_prompt_ids_event.set() + async with req_status.lock: + req_status.prefill_prompt_ids_event.prompt_ids = prompt_ids + req_status.prefill_prompt_ids_event.set() except: logger.error( f"PD_UPLOAD_PREFILL_PROMPT_IDS fail find req status for group_req_id: {group_req_id}" @@ -838,6 +827,7 @@ async def handle_loop(self): class ReqStatus: def __init__(self, req_id, p_node, d_node) -> None: self.req_id = req_id + self.lock = asyncio.Lock() self.event = asyncio.Event() self.up_status_event = asyncio.Event() self.prefill_prompt_ids_event = asyncio.Event() @@ -854,12 +844,14 @@ async def wait_to_ready(self): pass async def set_error(self, error_info: str, is_server_busy: bool = False): - # handle_loop and request consumers run on the same event loop. - self.error_info = error_info - self.is_server_busy = is_server_busy - self.event.set() - self.up_status_event.set() - self.prefill_prompt_ids_event.set() + async with self.lock: + self.error_info = error_info + self.is_server_busy = is_server_busy + # 请求可能正在等待 Prefill prompt ids、Decode KV 资源或输出 token, + # 设置全部事件,让请求自己的执行循环立即醒来并抛出异常。 + self.event.set() + self.up_status_event.set() + self.prefill_prompt_ids_event.set() def raise_if_error(self): if self.error_info is not None: @@ -871,26 +863,30 @@ def raise_if_error(self): ) raise RuntimeError(f"PD node generate failed: {self.error_info}") - def append_token(self, token_info: Tuple[int, str, dict, FinishStatus]): - # TOKEN_PACKS handling and fetch_pd_stream run on the same event loop. Keeping - # the mutation free of awaits makes the empty -> ready transition atomic. - was_empty = not self.out_token_info_list - self.out_token_info_list.append(token_info) - if was_empty: - self.event.set() + async def can_read(self, req_id_to_out_inf): + async with self.lock: + self.event.clear() + assert self.req_id in req_id_to_out_inf, f"error state req_id {self.req_id}" + if len(self.out_token_info_list) == 0: + return False + else: + return True - def pop_all_tokens(self): - self.event.clear() - ans = self.out_token_info_list - self.out_token_info_list = [] + async def pop_all_tokens(self): + async with self.lock: + ans = self.out_token_info_list.copy() + self.out_token_info_list.clear() return ans - def put_tokens_to_front(self, token_list: List[Tuple[int, str, dict, FinishStatus]]): + async def put_tokens_to_front(self, token_list: List[Tuple[int, str, dict, FinishStatus]]): if not token_list: return - self.out_token_info_list = token_list + self.out_token_info_list - self.event.set() + async with self.lock: + merged_tokens = token_list + self.out_token_info_list + self.out_token_info_list.clear() + self.out_token_info_list.extend(merged_tokens) + self.event.set() class PDManager: diff --git a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py index b8f8d6a1f2..a818fa5791 100644 --- a/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py +++ b/lightllm/server/httpserver_for_pd_master/pd_selector/cache_aware.py @@ -55,7 +55,7 @@ class CacheAwareConfig: # 每隔 sample_stride 个字符抽 1 个作为前缀树 key,降低匹配开销与内存。 sample_stride: int = 512 # 初始化前缀树时通过 sys.setrecursionlimit 调大 Python 调用栈深度。 - recursion_limit: int = 4000 + recursion_limit: int = field(default_factory=get_pd_master_recursion_limit) class BalanceRelThresholdController: diff --git a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py index 07a0be319b..542ddbd877 100644 --- a/lightllm/server/multi_level_kv_cache/disk_cache_worker.py +++ b/lightllm/server/multi_level_kv_cache/disk_cache_worker.py @@ -50,9 +50,8 @@ def __init__( # 读写同时进行时,分配8线程用来写,16线程用来读 max_concurrent_write_tasks = 8 - if disk_cache_dir: - cache_dir = os.path.join(disk_cache_dir, f"lightllm_disk_cache_{get_unique_server_name()}") - else: + cache_dir = disk_cache_dir + if not cache_dir: cache_dir = os.path.join(tempfile.gettempdir(), f"lightllm_disk_cache_{get_unique_server_name()}") os.makedirs(cache_dir, exist_ok=True) cache_file = os.path.join(cache_dir, "cache_file") diff --git a/lightllm/server/multimodal_params.py b/lightllm/server/multimodal_params.py index 8787e3ca8e..c3188f7f30 100644 --- a/lightllm/server/multimodal_params.py +++ b/lightllm/server/multimodal_params.py @@ -1,5 +1,4 @@ """Multimodal parameters for text generation.""" - import asyncio import os import librosa diff --git a/lightllm/server/pd_io_struct.py b/lightllm/server/pd_io_struct.py index bf39bcb570..7eff53c141 100644 --- a/lightllm/server/pd_io_struct.py +++ b/lightllm/server/pd_io_struct.py @@ -2,7 +2,7 @@ import time import copy from dataclasses import dataclass, field -from typing import Dict, List, Optional, Tuple +from typing import Dict, List, Optional from lightllm.server.req_id_generator import convert_sub_id_to_group_id from fastapi import WebSocket @@ -41,107 +41,8 @@ class ObjType(enum.Enum): PD_UPLOAD_PREFILL_PROMPT_IDS = 4 # prefill 节点上报生成的 prompt ids 信息。 PD_REQ_DECODE_NODE_INFO = 5 # pd master 节点下发给 prefill 节点的请求对应的 decode 节点信息。 HEARTBEAT = 6 # P/D 节点向 pd master 上报的心跳。 - TOKEN_PACKS_COMPACT = 7 # 不含 logprobs 等可选字段的紧凑 token 包。 - PD_UPLOAD_GENERATE_ERROR = 8 # P/D 节点向 pd master 上报本地请求生成异常。 - PD_UPLOAD_SERVER_BUSY = 9 # P/D 节点向 pd master 上报本地服务繁忙。 - - -PD_COMPACT_TOKEN_INFO_LEN = 12 -PDCompactTokenInfo = Tuple[ - int, # sub request id - str, # decoded text - int, # count_output_tokens - int, # prompt_tokens - int, # prompt_cache_len - int, # mtp_accepted_token_num - int, # mtp_verify_token_num - int, # mtp_verify_step_num - int, # finish status - Optional[str], # node_mode, first token only - Optional[Tuple[int, int, int]], # input text/audio/image tokens, first token only - Optional[int], # token id -] -_PD_COMPACT_METADATA_KEYS = frozenset( - { - "count_output_tokens", - "id", - "prompt_tokens", - "prompt_cache_len", - "mtp_accepted_token_num", - "mtp_verify_token_num", - "mtp_verify_step_num", - "node_mode", - "input_usage", - } -) -_PD_INPUT_USAGE_KEYS = frozenset({"input_text_tokens", "input_audio_tokens", "input_image_tokens"}) - - -def build_pd_compact_token_info(sub_req_id, text, metadata, finish_status) -> Optional[PDCompactTokenInfo]: - """Build the lossless compact form, or return None for optional metadata.""" - if not metadata.keys() <= _PD_COMPACT_METADATA_KEYS: - return None - - input_usage = metadata.get("input_usage") - compact_input_usage = None - if input_usage is not None: - if input_usage.keys() != _PD_INPUT_USAGE_KEYS: - return None - compact_input_usage = ( - input_usage["input_text_tokens"], - input_usage["input_audio_tokens"], - input_usage["input_image_tokens"], - ) - - return ( - sub_req_id, - text, - metadata["count_output_tokens"], - metadata["prompt_tokens"], - metadata["prompt_cache_len"], - metadata["mtp_accepted_token_num"], - metadata["mtp_verify_token_num"], - metadata["mtp_verify_step_num"], - finish_status.status, - metadata.get("node_mode"), - compact_input_usage, - metadata.get("id"), - ) - - -def unpack_pd_compact_token_info(token_info: PDCompactTokenInfo): - ( - sub_req_id, - text, - count_output_tokens, - prompt_tokens, - prompt_cache_len, - mtp_accepted_token_num, - mtp_verify_token_num, - mtp_verify_step_num, - finish_status, - node_mode, - input_usage, - token_id, - ) = token_info - metadata = { - "id": token_id, - "count_output_tokens": count_output_tokens, - "prompt_tokens": prompt_tokens, - "prompt_cache_len": prompt_cache_len, - "mtp_accepted_token_num": mtp_accepted_token_num, - "mtp_verify_token_num": mtp_verify_token_num, - "mtp_verify_step_num": mtp_verify_step_num, - } - if node_mode is not None: - metadata["node_mode"] = node_mode - if input_usage is not None: - metadata["input_usage"] = { - "input_text_tokens": input_usage[0], - "input_audio_tokens": input_usage[1], - "input_image_tokens": input_usage[2], - } - return sub_req_id, text, metadata, finish_status + PD_UPLOAD_GENERATE_ERROR = 7 # P/D 节点向 pd master 上报本地请求生成异常。 + PD_UPLOAD_SERVER_BUSY = 8 # P/D 节点向 pd master 上报本地服务繁忙。 @dataclass diff --git a/lightllm/server/router/manager.py b/lightllm/server/router/manager.py index 8ce4c63fac..43d03179c3 100644 --- a/lightllm/server/router/manager.py +++ b/lightllm/server/router/manager.py @@ -32,6 +32,7 @@ from lightllm.utils.graceful_utils import graceful_registry from lightllm.utils.process_check import start_parent_check_thread from lightllm.utils.envs_utils import get_unique_server_name +from lightllm.utils.config_utils import get_running_max_req_size_per_dp from lightllm.utils.shm_port_args import get_shm_port_args from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt from .stats import RouterStatics @@ -148,11 +149,7 @@ async def wait_to_model_ready(self): "weight_dir": self.model_weightdir, "load_way": self.load_way, "max_total_token_num": self.max_total_token_num, - "max_req_num": ( - self.args.running_max_req_size - if self.args.run_mode in ["prefill", "normal"] and self.args.enable_dp_prompt_cache_fetch - else self.args.per_dp_running_max_req_size - ), + "max_req_num": get_running_max_req_size_per_dp(self.args), # MTP length stopping is asynchronous, so up to mtp_step accepted # positions may already be committed when FINISHED_LENGTH is observed. # The overlapped iteration then needs mtp_step positions for target 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 ac254785ba..0f17ab664e 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -45,7 +45,10 @@ from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.common.basemodel.triton_kernel.gather_token_id import scatter_token from lightllm.server.pd_io_struct import PDChunckedTransTaskRet -from lightllm.server.multi_level_kv_cache import create_cache_placement_controller +from lightllm.server.multi_level_kv_cache import ( + CacheTier, + create_cache_placement_controller, +) from .multi_level_kv_cache import MultiLevelKvCacheModule from .dsv4_multi_level_kv_cache import Dsv4MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd @@ -199,11 +202,15 @@ def init_model(self, kvargs): ) # 初始化 dp 模式使用的通信 tensor, 对于非dp模式,不会使用到 if self.dp_size > 1: - self.dp_control_tensor = torch.zeros(2, dtype=torch.int32, device="cpu", requires_grad=False) self.dp_reduce_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) + self.dp_gather_item_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) + self.dp_all_gather_tensor = torch.tensor( + [0 for _ in range(self.global_world_size)], dtype=torch.int32, device="cuda", requires_grad=False + ) # 用于协同读取 ShmObjsIOBuffer 中的请求信息的通信tensor和通信组对象。 - self.node_broadcast_tensor = torch.zeros(1, dtype=torch.int32, device="cpu", requires_grad=False) + self.node_broadcast_tensor = torch.tensor([0], dtype=torch.int32, device="cuda", requires_grad=False) + # DeepSeek-V4 DP prompt-cache checkpoints live in process-local CPU memory. self.node_gloo_group = create_new_group_for_current_node("gloo") self.node_nccl_group = create_new_group_for_current_node("nccl") @@ -472,8 +479,8 @@ def _try_read_new_reqs_normal(self): self.node_broadcast_tensor.fill_(0) src_rank_id = self.args.node_rank * self.node_world_size - broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_gloo_group, async_op=False) - new_buffer_is_ready = self.node_broadcast_tensor.item() + broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_nccl_group, async_op=False) + new_buffer_is_ready = self.node_broadcast_tensor.detach().item() if new_buffer_is_ready: self._read_reqs_buffer_and_init_reqs() @@ -486,8 +493,8 @@ def _try_read_new_reqs_normal(self): self.node_broadcast_tensor.fill_(0) src_rank_id = self.args.node_rank * self.node_world_size - broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_gloo_group, async_op=False) - new_buffer_is_ready = self.node_broadcast_tensor.item() + broadcast(self.node_broadcast_tensor, src=src_rank_id, group=self.node_nccl_group, async_op=False) + new_buffer_is_ready = self.node_broadcast_tensor.detach().item() if new_buffer_is_ready: self._read_pd_trans_io_buffer_and_update_req_status() return @@ -823,14 +830,24 @@ def _get_classed_reqs( req_obj.wait_pause = True wait_pause_count += 1 - # 先由控制器确定请求需要写入的缓存层级。 + # 先由控制器确定请求需要写入的缓存层级,再按是否包含 CPU cache 决定是否发起 offload。 cache_controller = g_infer_context.cache_placement_controller new_finished_reqs = [req for req in finished_reqs if req.cpu_cache_task_status.is_not_started()] cache_controller.set_req_cache_way(new_finished_reqs) if self.args.enable_cpu_cache: - true_finished_reqs = self.multi_level_cache_module.offload_finished_reqs_to_cpu_cache( - finished_reqs=finished_reqs + offload_reqs = [ + req for req in finished_reqs if CacheTier.CPU in req.cache_tiers or CacheTier.DISK in req.cache_tiers + ] + offload_finished_reqs = self.multi_level_cache_module.offload_finished_reqs_to_cpu_cache( + finished_reqs=offload_reqs ) + offload_finished_req_ids = {req.req_id for req in offload_finished_reqs} + true_finished_reqs = [ + req + for req in finished_reqs + if (CacheTier.CPU not in req.cache_tiers and CacheTier.DISK not in req.cache_tiers) + or req.req_id in offload_finished_req_ids + ] else: true_finished_reqs = finished_reqs @@ -1008,26 +1025,23 @@ def _sample_and_scatter_token( ) return next_token_ids, next_token_ids_cpu, next_token_logprobs_cpu, next_token_ranks_cpu - def _dp_all_reduce_req_presence( + def _dp_all_gather_prefill_and_decode_req_num( self, prefill_reqs: List[InferReq], decode_reqs: List[InferReq] - ) -> tuple[bool, bool]: + ) -> Tuple[np.ndarray, np.ndarray]: """ - Return whether any DP rank has prefill or decode requests. - - Request counts originate on the CPU and the scheduler only needs their - global presence. Keep this control-plane collective on the CPU so it can - overlap the previous CUDA graph instead of synchronizing that graph back - to the host every decode step. + Gather the number of prefill requests across all DP ranks. """ - self.dp_control_tensor[0] = bool(prefill_reqs) - self.dp_control_tensor[1] = bool(decode_reqs) - all_reduce( - self.dp_control_tensor, - op=dist.ReduceOp.MAX, - group=dist_group_manager.dp_control_group, - async_op=False, - ) - return bool(self.dp_control_tensor[0]), bool(self.dp_control_tensor[1]) + current_dp_prefill_num = len(prefill_reqs) + self.dp_gather_item_tensor.fill_(current_dp_prefill_num) + all_gather_into_tensor(self.dp_all_gather_tensor, self.dp_gather_item_tensor, group=None, async_op=False) + dp_prefill_req_nums = self.dp_all_gather_tensor.cpu().numpy() + + current_dp_decode_num = len(decode_reqs) + self.dp_gather_item_tensor.fill_(current_dp_decode_num) + all_gather_into_tensor(self.dp_all_gather_tensor, self.dp_gather_item_tensor, group=None, async_op=False) + dp_decode_req_nums = self.dp_all_gather_tensor.cpu().numpy() + + return dp_prefill_req_nums, dp_decode_req_nums def _dp_all_reduce_decode_req_num(self, decode_reqs: List[InferReq]) -> int: """ diff --git a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py index 683d345d4d..b159b98f25 100644 --- a/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py +++ b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl_for_xgrammar_mode.py @@ -4,7 +4,6 @@ from .impl import ChunkedPrefillBackend from lightllm.utils.infer_utils import calculate_time from lightllm.server.core.objs import FinishStatus -from lightllm.server.core.objs.sampling_params import get_xgrammar_tokenizer from lightllm.server.router.model_infer.infer_batch import g_infer_context, InferReq from lightllm.server.tokenizer import get_tokenizer from lightllm.utils.log_utils import init_logger @@ -27,7 +26,7 @@ def init_custom(self): self.args.model_dir, self.args.tokenizer_mode, trust_remote_code=self.args.trust_remote_code ) - self.tokenizer_info = xgr.TokenizerInfo.from_huggingface(get_xgrammar_tokenizer(self.tokenizer)) + self.tokenizer_info = xgr.TokenizerInfo.from_huggingface(self.tokenizer) self.xgrammar_compiler = xgr.GrammarCompiler(self.tokenizer_info, max_threads=8) self.xgrammar_token_bitmask = xgr.allocate_token_bitmask(1, self.tokenizer_info.vocab_size) diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py index c90a864731..ffd92202d8 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/control_state.py @@ -1,5 +1,8 @@ +import numpy as np from enum import Enum +from typing import List from lightllm.utils.envs_utils import get_env_start_args +from lightllm.server.router.model_infer.infer_batch import InferReq from ..base_backend import ModeBackend @@ -17,8 +20,10 @@ def __init__(self, backend: ModeBackend): def select_run_way( self, - has_prefill: bool, - has_decode: bool, + dp_prefill_req_nums: np.ndarray, + dp_decode_req_nums: np.ndarray, + prefill_reqs: List[InferReq], + decode_reqs: List[InferReq], ) -> "RunWay": """ 判断决策运行方式: @@ -27,41 +32,55 @@ def select_run_way( self.step_count += 1 if self.is_aggressive_schedule: return self._agressive_way( - has_prefill=has_prefill, - has_decode=has_decode, + dp_prefill_req_nums=dp_prefill_req_nums, + dp_decode_req_nums=dp_decode_req_nums, + prefill_reqs=prefill_reqs, + decode_reqs=decode_reqs, ) else: return self._normal_way( - has_prefill=has_prefill, - has_decode=has_decode, + dp_prefill_req_nums=dp_prefill_req_nums, + dp_decode_req_nums=dp_decode_req_nums, + prefill_reqs=prefill_reqs, + decode_reqs=decode_reqs, ) def _agressive_way( self, - has_prefill: bool, - has_decode: bool, + dp_prefill_req_nums: np.ndarray, + dp_decode_req_nums: np.ndarray, + prefill_reqs: List[InferReq], + decode_reqs: List[InferReq], ): - if has_prefill: + max_prefill_num = np.max(dp_prefill_req_nums) + if max_prefill_num > 0: return RunWay.PREFILL - if has_decode: + max_decode_num = np.max(dp_decode_req_nums) + if max_decode_num > 0: return RunWay.DECODE return RunWay.PASS def _normal_way( self, - has_prefill: bool, - has_decode: bool, + dp_prefill_req_nums: np.ndarray, + dp_decode_req_nums: np.ndarray, + prefill_reqs: List[InferReq], + decode_reqs: List[InferReq], ): - if self.left_decode_num > 0 and has_decode: + # use_ratio = np.count_nonzero(dp_prefill_req_nums) / dp_prefill_req_nums.shape[0] + max_decode_num = np.max(dp_decode_req_nums) + max_prefill_num = np.max(dp_prefill_req_nums) + + if self.left_decode_num > 0 and max_decode_num > 0: self.left_decode_num -= 1 return RunWay.DECODE - if has_prefill: + if max_prefill_num > 0: # prefill 一次允许进行几次 decode 操作。 self.left_decode_num = self.decode_max_step return RunWay.PREFILL else: - if has_decode: + if max_decode_num > 0: return RunWay.DECODE else: return RunWay.PASS diff --git a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py index 8be392e17c..86efda9586 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py @@ -298,7 +298,7 @@ def kv_trans(self, trans_tasks: List["TransTask"]): checkpoint_len=end, ) - if self.backend.is_deepseek_v4: + if self.backend.is_deepseek_v4 and self.backend.args.enable_cpu_cache: # CPU-cache restore can evict source radix pages before the scheduler all-gather fences this stream. dist.barrier(group=self.backend.node_nccl_group) 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 304eede019..13f9feee7d 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 @@ -135,13 +135,15 @@ def infer_loop(self): recover_paused=self.control_state_machine.try_recover_paused_reqs(), ) - has_prefill, has_decode = self._dp_all_reduce_req_presence( + dp_prefill_req_nums, dp_decode_req_nums = self._dp_all_gather_prefill_and_decode_req_num( prefill_reqs=prefill_reqs, decode_reqs=decode_reqs ) run_way = self.control_state_machine.select_run_way( - has_prefill=has_prefill, - has_decode=has_decode, + dp_prefill_req_nums=dp_prefill_req_nums, + dp_decode_req_nums=dp_decode_req_nums, + prefill_reqs=prefill_reqs, + decode_reqs=decode_reqs, ) if run_way.is_prefill(): diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py index 2af353f796..d6a02fa19e 100644 --- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py @@ -167,10 +167,6 @@ def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> L true_finished_reqs = [] cpu_stream = g_infer_context.get_cpu_kv_cache_stream() for req in finished_reqs: - if CacheTier.CPU not in req.cache_tiers: - true_finished_reqs.append(req) - continue - # 只有 group_req_id 和 request_id 相同的请求才会被卸载到 cpu cache 中。 # 这个限制是为了兼容 diverse 模式下的请求处理, 只有主请求才 offload kv 到 cpu # cache 中 diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index f8f0045fb9..e076cce8a0 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -69,7 +69,7 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: req_obj: InferReq = g_infer_context.requests_mapping[request_id] # pending 期间优先重新匹配 radix;准入失败时释放引用,留待下轮重试。 - if req_obj.pd_task_num == 0 and not req_obj.infer_aborted: + if self.is_deepseek_v4 and req_obj.pd_task_num == 0 and not req_obj.infer_aborted: if g_infer_context.is_hybrid_att_model: req_obj._hybrid_match_radix_cache() else: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py index f015264a8b..09e2d36364 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_trans_process.py @@ -300,8 +300,7 @@ def accept_peer_task_loop( local_trans_task.prefill_page_reg_desc = remote_trans_task.prefill_page_reg_desc local_trans_task.transfer_nbytes = remote_trans_task.transfer_nbytes self.request_page_task_queue.put(local_trans_task) - if self.args.detail_log: - logger.info(f"recv WRITE request from prefill: {remote_trans_task.to_str()}") + logger.info(f"recv WRITE request from prefill: {remote_trans_task.to_str()}") else: # This does not necessarily mean the WRITE protocol state is corrupted. # A common benign case is: decode has already received an abort for this @@ -328,8 +327,7 @@ def accept_peer_task_loop( local_trans_task.first_gen_token_id = remote_trans_task.first_gen_token_id local_trans_task.first_gen_token_logprob = remote_trans_task.first_gen_token_logprob self.ready_page_task_queue.put(local_trans_task) - if self.args.detail_log: - logger.info(f"recv WRITE done from prefill: {remote_trans_task.to_str()}") + logger.info(f"recv WRITE done from prefill: {remote_trans_task.to_str()}") else: # Same race as the WRITE request stage: decode may have cleaned the # waiting task because the request was aborted, then a late done notify @@ -432,14 +430,13 @@ def success_loop(self): ret = trans_task.createRetObj() self.task_out_queue.put(ret) - if self.args.detail_log: - if trans_task.start_trans_time is not None: - logger.info( - f"trans task ret success:{ret} cost time: {trans_task.transfer_time()} s " - f"read_page_gpu_time: {read_page_gpu_time_ms:.3f} ms" - ) - else: - logger.info(f"trans task ret success:{ret}") + if trans_task.start_trans_time is not None: + logger.info( + f"trans task ret success:{ret} cost time: {trans_task.transfer_time()} s " + f"read_page_gpu_time: {read_page_gpu_time_ms:.3f} ms" + ) + else: + logger.info(f"trans task ret success:{ret}") @log_exception def fail_loop(self): diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index c2b0d94691..cbfcc1ee4f 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -16,7 +16,6 @@ from nixl._api import nixl_agent as NixlWrapper from nixl._api import nixlBind from nixl._api import nixl_agent_config - from nixl._api import nixl_thread_sync_t logger.info("Nixl is available") except ImportError: @@ -34,13 +33,13 @@ def __init__(self, node_id: int, tp_idx: int, kv_move_buffer: Tensor): "yes", "on", ) - conf = nixl_agent_config(sync_mode=nixl_thread_sync_t.NIXL_THREAD_SYNC_RW) + conf = None if self.capture_telemetry: + conf = nixl_agent_config() conf.capture_telemetry = True logger.info("NIXL telemetry enabled") self.nixl_agent = NixlWrapper(self.agent_name, conf) self._register_kv_move_buffer(kv_move_buffer=kv_move_buffer) - self._remote_agents_lock = threading.Lock() self.remote_agents: Dict[str, PDAgentMetadata] = {} # Serialize complete peer add/remove operations, including native NIXL # calls and descriptor cleanup, across worker threads (see GH-1470). @@ -129,8 +128,7 @@ def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): } logger.info( - f"Added remote agent {peer_name} with mem desc {page_mem_desc} " - f"cost time: {time.time() - start_time} s" + f"Added remote agent {peer_name} with mem desc {page_mem_desc} cost time: {time.time() - start_time} s" ) self.remote_agents[remote_agent.agent_name] = remote_agent @@ -314,12 +312,7 @@ def write_blocks_paged( def check_task_status(self, trans_task: PDChunckedTransTask) -> str: assert trans_task.xfer_handle is not None handle = trans_task.xfer_handle - try: - xfer_state = self.nixl_agent.check_xfer_state(handle) - except Exception as e: - logger.error(f"Check transfer state failed with trans task {trans_task.to_str()} for handle {handle}") - logger.exception(str(e)) - return "ERR" + xfer_state = self.nixl_agent.check_xfer_state(handle) if xfer_state == "ERR": logger.warning(f"Transfer failed with trans task {trans_task.to_str()} for handle {handle}") return xfer_state diff --git a/lightllm/server/router/req_queue/base_queue.py b/lightllm/server/router/req_queue/base_queue.py index 0d28521091..0c25b8949a 100644 --- a/lightllm/server/router/req_queue/base_queue.py +++ b/lightllm/server/router/req_queue/base_queue.py @@ -3,7 +3,7 @@ from lightllm.utils.infer_utils import calculate_time from ..batch import Batch, Req from lightllm.server.core.objs import FinishStatus -from lightllm.utils.config_utils import get_fixed_kv_len +from lightllm.utils.config_utils import get_fixed_kv_len, get_running_max_req_size_per_dp from lightllm.server.core.objs import StartArgs from lightllm.utils.log_utils import init_logger @@ -23,11 +23,7 @@ def __init__(self, args: StartArgs, router, dp_index, dp_size_in_node) -> None: # 在极端情况下减少,在非特定模式下,get_fixed_kv_len() 返回的都是 # 0, 不会有任何影响。 self.max_total_tokens = args.max_total_token_num - get_fixed_kv_len() - assert args.batch_max_tokens is not None - self.batch_max_tokens = args.batch_max_tokens - if args.per_dp_running_max_req_size is None: - raise RuntimeError("per_dp_running_max_req_size is not initialized") - self.running_max_req_size = args.per_dp_running_max_req_size + self.running_max_req_size = get_running_max_req_size_per_dp(args) self.waiting_req_list: List[Req] = [] # List of queued requests self.router_token_ratio = args.router_token_ratio # ratio to determine whether the router is busy diff --git a/lightllm/server/router/req_queue/dp_balancer/__init__.py b/lightllm/server/router/req_queue/dp_balancer/__init__.py index 8271236899..34f994f8a2 100644 --- a/lightllm/server/router/req_queue/dp_balancer/__init__.py +++ b/lightllm/server/router/req_queue/dp_balancer/__init__.py @@ -2,7 +2,6 @@ from typing import List from lightllm.server.router.req_queue.base_queue import BaseQueue from .bs import DpBsBalancer -from .cache_aware import DpCacheAwareBalancer def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): @@ -10,9 +9,5 @@ def get_dp_balancer(args, dp_size_in_node: int, inner_queues: List[BaseQueue]): return RoundRobinDpBalancer(dp_size_in_node, inner_queues) elif args.dp_balancer == "bs_balancer": return DpBsBalancer(dp_size_in_node, inner_queues) - elif args.dp_balancer == "cache_aware": - if args.disable_dynamic_prompt_cache: - raise ValueError("cache_aware DP balancing requires dynamic prompt cache") - return DpCacheAwareBalancer(dp_size_in_node, inner_queues, run_mode=args.run_mode) else: raise ValueError(f"Invalid dp balancer: {args.dp_balancer}") diff --git a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py b/lightllm/server/router/req_queue/dp_balancer/cache_aware.py deleted file mode 100644 index 6a59270624..0000000000 --- a/lightllm/server/router/req_queue/dp_balancer/cache_aware.py +++ /dev/null @@ -1,188 +0,0 @@ -"""DP-local cache-affinity routing based on bounded token-prefix history. - -The router owns this heuristic index. It records dispatch history rather than querying -the infer processes' radix trees, so stale entries can only affect placement, not KV -cache correctness. Prefill load uses remaining prompt tokens; other modes count requests. -""" - -from __future__ import annotations - -import random -from collections import OrderedDict -from dataclasses import dataclass -from typing import List, Optional, Tuple - -import xxhash - -from lightllm.server.router.batch import Batch, Req -from lightllm.server.router.req_queue.base_queue import BaseQueue - -from .base import DpBalancer - - -PrefixHash = Tuple[int, int] - - -@dataclass(slots=True) -class DpCacheAwareConfig: - cache_threshold: float = 0.5 - balance_rel_threshold: float = 1.8 - # DeepSeek-V4 prompt-cache entries are reusable only at 256-token boundaries. - block_size: int = 256 - max_cache_entries: int = 1_000_000 - evict_entries: int = 10_000 - - -class TokenPrefixCache: - """Bounded LRU mapping from cumulative token-prefix hashes to local DP indexes.""" - - def __init__(self, block_size: int, max_entries: int, evict_entries: int) -> None: - if block_size < 1: - raise ValueError(f"block_size must be >= 1, got {block_size}") - if max_entries < 0: - raise ValueError(f"max_entries must be >= 0, got {max_entries}") - if evict_entries < 1: - raise ValueError(f"evict_entries must be >= 1, got {evict_entries}") - self.block_size = block_size - self.max_entries = max_entries - self.evict_entries = evict_entries - self._prefix_to_dp: OrderedDict[int, int] = OrderedDict() - - def hash_prefixes(self, prompt_ids) -> List[PrefixHash]: - cacheable_token_count = max(0, len(prompt_ids) - 1) - cacheable_token_count = cacheable_token_count // self.block_size * self.block_size - if cacheable_token_count == 0: - return [] - - token_view = memoryview(prompt_ids) - item_size = token_view.itemsize - prompt_bytes = token_view.cast("B") - # A collision only changes a routing hint; it cannot affect cache correctness. - # xxh3-64 keeps the 1M-entry index compact and hashes 1M-token prompts faster. - hasher = xxhash.xxh3_64() - prefix_hashes = [] - for start in range(0, cacheable_token_count, self.block_size): - end = start + self.block_size - hasher.update(prompt_bytes[start * item_size : end * item_size]) - prefix_hashes.append((hasher.intdigest(), end)) - return prefix_hashes - - def match(self, prefix_hashes: List[PrefixHash]) -> Tuple[Optional[int], int]: - for prefix_hash, token_count in reversed(prefix_hashes): - try: - dp_index = self._prefix_to_dp[prefix_hash] - except KeyError: - continue - self._prefix_to_dp.move_to_end(prefix_hash) - return dp_index, token_count - return None, 0 - - def insert(self, prefix_hashes: List[PrefixHash], dp_index: int, start_index: int = 0) -> None: - for prefix_index in range(start_index, len(prefix_hashes)): - prefix_hash = prefix_hashes[prefix_index][0] - self._prefix_to_dp[prefix_hash] = dp_index - self._prefix_to_dp.move_to_end(prefix_hash) - - if len(self._prefix_to_dp) > self.max_entries: - evict_count = len(self._prefix_to_dp) - self.max_entries + self.evict_entries - for _ in range(min(evict_count, len(self._prefix_to_dp))): - self._prefix_to_dp.popitem(last=False) - - def __len__(self) -> int: - return len(self._prefix_to_dp) - - -class DpCacheAwareBalancer(DpBalancer): - """Route matching token prefixes to the same local DP unless load requires rebalancing.""" - - def __init__( - self, - dp_size_in_node: int, - inner_queues: List[BaseQueue], - run_mode: str, - config: Optional[DpCacheAwareConfig] = None, - ) -> None: - super().__init__(dp_size_in_node, inner_queues) - self.run_mode = run_mode - self.config = config or DpCacheAwareConfig() - self.prefix_cache = TokenPrefixCache( - block_size=self.config.block_size, - max_entries=self.config.max_cache_entries, - evict_entries=self.config.evict_entries, - ) - - def assign_reqs_to_dp(self, current_batch: Batch, reqs_waiting_for_dp_index: List[List[Req]]) -> None: - if not reqs_waiting_for_dp_index: - return - - if self.run_mode == "prefill": - # Queued requests have not matched real KV yet; dispatch history is only a routing hint. - total_load_per_dp = [sum(req.input_len for req in queue.waiting_req_list) for queue in self.inner_queues] - if current_batch is not None: - for req in current_batch.reqs: - total_load_per_dp[req.sample_params.suggested_dp_index] += max( - 0, req.input_len - req.shm_cur_kv_len - ) - else: - current_load_per_dp = [0 for _ in range(self.dp_size_in_node)] - if current_batch is not None: - current_load_per_dp = current_batch.get_all_dp_req_num() - total_load_per_dp = [ - current_load_per_dp[dp_index] + len(self.inner_queues[dp_index].waiting_req_list) - for dp_index in range(self.dp_size_in_node) - ] - - for req_group in reqs_waiting_for_dp_index: - first_req = req_group[0] - group_load = sum(req.input_len for req in req_group) if self.run_mode == "prefill" else len(req_group) - linked_prompt_ids = False - if not hasattr(first_req, "shm_prompt_ids"): - first_req.link_prompt_ids_shm_array() - linked_prompt_ids = True - try: - prefix_hashes = self.prefix_cache.hash_prefixes(first_req.get_prompt_ids_numpy()) - finally: - if linked_prompt_ids: - first_req.shm_prompt_ids.detach_shm() - del first_req.shm_prompt_ids - - cache_dp_index = None - matched_token_count = 0 - if not first_req.sample_params.disable_prompt_cache: - matched_dp_index, matched_token_count = self.prefix_cache.match(prefix_hashes) - match_rate = matched_token_count / first_req.input_len if first_req.input_len else 0.0 - if match_rate > self.config.cache_threshold: - cache_dp_index = matched_dp_index - - idle_dp_indexes = [dp_index for dp_index, load in enumerate(total_load_per_dp) if load == 0] - if idle_dp_indexes: - if cache_dp_index in idle_dp_indexes: - selected_dp_index = cache_dp_index - else: - selected_dp_index = random.choice(idle_dp_indexes) - else: - min_load = min(total_load_per_dp) - least_loaded_dp_indexes = [ - dp_index for dp_index, load in enumerate(total_load_per_dp) if load == min_load - ] - least_loaded_dp_index = random.choice(least_loaded_dp_indexes) - if cache_dp_index is None: - selected_dp_index = least_loaded_dp_index - else: - cache_projected_load = total_load_per_dp[cache_dp_index] + group_load - least_projected_load = total_load_per_dp[least_loaded_dp_index] + group_load - if cache_projected_load > least_projected_load * self.config.balance_rel_threshold: - selected_dp_index = least_loaded_dp_index - else: - selected_dp_index = cache_dp_index - - for req in req_group: - req.sample_params.suggested_dp_index = selected_dp_index - self.inner_queues[selected_dp_index].extend(req_group) - total_load_per_dp[selected_dp_index] += group_load - insert_start_index = 0 - if cache_dp_index == selected_dp_index: - insert_start_index = (matched_token_count + self.config.block_size - 1) // self.config.block_size - self.prefix_cache.insert(prefix_hashes, selected_dp_index, start_index=insert_start_index) - - reqs_waiting_for_dp_index.clear() diff --git a/lightllm/server/tokenizer.py b/lightllm/server/tokenizer.py index e372a1f925..cc7ecd8a8b 100644 --- a/lightllm/server/tokenizer.py +++ b/lightllm/server/tokenizer.py @@ -61,9 +61,15 @@ def get_tokenizer( try: tokenizer = AutoTokenizer.from_pretrained(tokenizer_name, trust_remote_code=trust_remote_code, *args, **kwargs) except ValueError as e: - if tokenizer_mode == "slow" or "Tokenizer class TokenizersBackend does not exist" not in str(e): + if ( + model_type != "deepseek_v4" + or tokenizer_mode == "slow" + or "Tokenizer class TokenizersBackend does not exist" not in str(e) + ): raise - logger.warning("Transformers does not provide TokenizersBackend; loading tokenizer.json as a fast tokenizer") + logger.warning( + "Transformers does not provide DeepSeek-V4 TokenizersBackend; loading tokenizer.json as a fast tokenizer" + ) tokenizer = PreTrainedTokenizerFast.from_pretrained(tokenizer_name, *args, **kwargs) except TypeError as e: # The LLaMA tokenizer causes a protobuf error in some environments, using slow mode. diff --git a/lightllm/server/visualserver/model_infer/model_rpc.py b/lightllm/server/visualserver/model_infer/model_rpc.py index 6f8ebed777..95209ef6bc 100644 --- a/lightllm/server/visualserver/model_infer/model_rpc.py +++ b/lightllm/server/visualserver/model_infer/model_rpc.py @@ -22,6 +22,7 @@ from lightllm.models.qwen3_vl.qwen3_visual import Qwen3VisionTransformerPretrainedModel from lightllm.models.tarsier2.tarsier2_visual import TarsierVisionTransformerPretrainedModel from lightllm.models.qwen3_omni_moe_thinker.qwen3_omni_visual import Qwen3OmniMoeVisionTransformerPretrainedModel +from lightllm.models.neo_chat_moe.neo_visual import NeoVisionTransformerPretrainedModel from lightllm.utils.infer_utils import set_random_seed from lightllm.utils.dist_utils import init_vision_distributed_env from lightllm.utils.envs_utils import get_env_start_args @@ -114,8 +115,6 @@ def exposed_init_model(self, kvargs): .bfloat16() ) elif self.model_type == "neo_chat": - from lightllm.models.neo_chat_moe.neo_visual import NeoVisionTransformerPretrainedModel - self.model = NeoVisionTransformerPretrainedModel(kvargs, **model_cfg["vision_config"]).eval().bfloat16() else: raise Exception(f"can not support {self.model_type} now") diff --git a/lightllm/utils/config_utils.py b/lightllm/utils/config_utils.py index bbda35427b..e9cd481ec5 100644 --- a/lightllm/utils/config_utils.py +++ b/lightllm/utils/config_utils.py @@ -81,11 +81,13 @@ def get_config_json(model_path: str): return normalize_deepseek_v4_config(json_obj) -def get_generation_config_diff_dict(model_path: str) -> Dict[str, Any]: +def get_generation_config_dict(model_path: str) -> Dict[str, Any]: from transformers import GenerationConfig - generation_cfg = GenerationConfig.from_pretrained(model_path, trust_remote_code=True).to_diff_dict() - return {key: value for key, value in generation_cfg.items() if value is not None} + generation_cfg = GenerationConfig.from_pretrained(model_path, trust_remote_code=True) + if get_model_type(model_path) == "deepseek_v4": + return {key: value for key, value in generation_cfg.to_diff_dict().items() if value is not None} + return generation_cfg.to_dict() def get_running_max_req_size_per_dp(args) -> int: @@ -647,8 +649,12 @@ def get_reasoning_parser_for_model(model_path: str) -> Optional[str]: ]: return "qwen3" - # DeepSeek V3 / V4 (share the ... reasoning format, request-gated) - if model_type in ["deepseek_v3", "deepseek_v31", "deepseek_v32", "deepseek_v4"]: + # DeepSeek V4 + if model_type == "deepseek_v4": + return "deepseek-v4" + + # DeepSeek V3 + if model_type in ["deepseek_v3", "deepseek_v31", "deepseek_v32"]: return "deepseek-v3" # DeepSeek R1 diff --git a/lightllm/utils/device_utils.py b/lightllm/utils/device_utils.py index 750bdbd4d6..58bff90560 100644 --- a/lightllm/utils/device_utils.py +++ b/lightllm/utils/device_utils.py @@ -145,9 +145,8 @@ def has_nvlink(): # Call nvidia-smi to get the topology matrix result = subprocess.check_output(["nvidia-smi", "topo", "--matrix"]) result = result.decode("utf-8") - # NVLink topology entries are reported as NV followed by the link count, - # for example NV8 on H800 and NV18 on B300. - return any(entry.startswith("NV") and entry[2:].isdigit() for entry in result.split()) + # Check if the output contains 'NVLink' + return any(f"NV{i}" in result for i in range(1, 8)) except FileNotFoundError: # nvidia-smi is not installed, assume no NVLink return False diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py index f4e4a41f13..c5b88d5799 100644 --- a/lightllm/utils/envs_utils.py +++ b/lightllm/utils/envs_utils.py @@ -37,31 +37,6 @@ def set_env_start_args(args): if not isinstance(args, dict): args = vars(args) os.environ["LIGHTLLM_START_ARGS"] = json.dumps(args) - if args["enable_ep_moe"]: - if args["run_mode"] == "prefill": - decode_capacity = args["running_max_req_size"] * (args["mtp_step"] + 1) - decode_capacity = ((decode_capacity + 7) // 8) * 8 - configured_decode_capacity = int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", decode_capacity)) - if configured_decode_capacity != decode_capacity: - logger.warning( - "NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE=%d differs from the automatically derived value %d.", - configured_decode_capacity, - decode_capacity, - ) - decode_capacity = max(configured_decode_capacity, decode_capacity) - else: - decode_capacity = get_deepep_num_max_dispatch_tokens_per_rank_decode() - min_qp_depth = 2 * (decode_capacity + 1) - # NVSHMEM IBGDA rejects QP depths below NVSHMEMI_IBGDA_MIN_QP_DEPTH. - derived_qp_depth = max(128, 1 << (min_qp_depth - 1).bit_length()) - configured_qp_depth = int(os.getenv("NVSHMEM_QP_DEPTH", derived_qp_depth)) - if configured_qp_depth < derived_qp_depth: - logger.warning( - "NVSHMEM_QP_DEPTH=%d is below the required minimum; using %d instead.", - configured_qp_depth, - derived_qp_depth, - ) - os.environ["NVSHMEM_QP_DEPTH"] = str(max(configured_qp_depth, derived_qp_depth)) return @@ -118,27 +93,8 @@ def get_deepep_num_max_dispatch_tokens_per_rank_prefill(): @lru_cache(maxsize=None) def get_deepep_num_max_dispatch_tokens_per_rank_decode(): - args = get_env_start_args() - per_dp_running_max_req_size = getattr(args, "per_dp_running_max_req_size", None) - if per_dp_running_max_req_size is None: - per_dp_running_max_req_size = args.running_max_req_size - - graph_max_batch_size = 0 - if not args.disable_cudagraph: - graph_max_batch_size = args.graph_max_batch_size - if args.enable_decode_microbatch_overlap: - graph_max_batch_size //= 2 - - required = max(per_dp_running_max_req_size, graph_max_batch_size) * (args.mtp_step + 1) - required = ((required + 7) // 8) * 8 - configured = int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", required)) - if configured != required: - logger.warning( - "NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE=%d differs from the automatically derived value %d.", - configured, - required, - ) - return max(configured, required) + # 该参数需要大于单卡最大batch size,且是8的倍数。该参数与显存占用直接相关,值越大,显存占用越大,如果出现显存不足,可以尝试调小该值 + return int(os.getenv("NUM_MAX_DISPATCH_TOKENS_PER_RANK_DECODE", 256)) @lru_cache(maxsize=None) diff --git a/revert.md b/revert.md index 20fd8d2cfe..4039c36036 100644 --- a/revert.md +++ b/revert.md @@ -45,3 +45,23 @@ | `ca4f075860e00649ea73554780b33a417ead21bc` (`synchronize PDL top-k and fail fast on model thread errors`) | DeepSeek-V4 PDL top-k 同步 | 通用 fatal thread excepthook | 整理阶段产生的临时 revert 提交已经压平,因此最终分支相对 `7369626` 只多一个汇总清理提交;本文件记录被排除的原始提交,便于后续追溯。 + +## 2026-09-28 二次模型边界清理 + +本节是在 `support_dsv4_model@c370e9213f1632423f4c9478258bdf377b18715c` 上追加的第二次清理记录,不改写上面的首次拆分历史。本轮仍保留 DeepSeek-V4 基础模型、MTP、DSpark、Vision、PD、多级缓存及其必要的共享层契约;专家能力和通用稳定性/运维继续按上文排除。 + +| 类别 | 回退内容 | 对应原提交或路径 | +| --- | --- | --- | +| 通用 DP 调度与容量架构 | 恢复既有 NCCL 调度控制、统一请求容量和通用 DP 状态;仅为 DSV4 跨 DP CPU checkpoint 保留节点内 Gloo group | `94fdfdeca05f3f94619a0c183088114c42e65124`、`f2e8f1429dccd199750bc5d9334dc4a30622dd15`、`c211f485aa0c5dcd392a22104f5b37cf44056c2e` | +| DP cache-aware 调度 | 删除新增的 router DP cache-aware balancer,并恢复 `req_queue`、router manager 和控制状态 | `fbe49a2227d714ce373a0fde9d5701c77368ed85` | +| 通用 HTTP/PD 传输优化 | 回退 compact token payload、streaming/serialization、async queue 和 PD-master 增量;保留 DSV4 packed cache 所需的变长字节传输 | `cfce08726b09e44182ff63df2978ae95cfea1337`、`baadd7ea316861c089956d46021e0b216a51bcf8`、`095c043d4781c8cc8f8b515d4e4e618958ff5a6f` | +| MXFP4 Marlin 运行时后端 | 删除 `mxfp4w4a16-b32-marlin` 注册、实现、权重 finalize 和 CLI 暴露;保留 DSV4 MTP checkpoint 转 BF16 工具 | `e8009cb3e053ffe7dbe465c027e4fe6a676181c8` 中的可选 Marlin 路径 | +| xgrammar 通用 tokenizer 兼容 | 删除 `get_xgrammar_tokenizer` 及底层 HF tokenizer 旁路;DSV4 tokenizer 回退只对 DSV4 生效 | `b8073f843c71dcd3f837b3b17522383dc3688b9d` | +| structured-output 通用降级 | 恢复原有约束输出行为,不在 xgrammar 缺失时静默关闭 | `f14dd40738daf817d66ea8d025a93fe154fa8a8b`、`8ee409a5610d5a4a1023942a07ead1f020d269ec` | +| Disk cache 实例目录隔离 | 恢复原有 disk worker/目录语义 | `0ed251199d7df166be65974e9846599409de94f8` | +| Benchmark 与 autotune 产物 | 删除 static benchmark 增量及本 PR 新增的 H100/H200 autotune JSON | `388dd29de8a000500279ad273c06a6d7874f7cf1`、`4ac53e82fd3ecabef8ff0bbd4100a5fec3202821`、`0a48e27bcbfa8fec91f52bebb874616e1a2a95a0` | +| DeepEP 环境自动调参 | 回退 decode dispatch capacity、NVSHMEM QP depth 等通用自动派生,恢复原有默认值 | `9a071ca5e0e53bab4bca63586f78e92b493cfb68`、`97ae2d12acdd77bcf87a786a710c0f1f049d8196`、`69c260f4a5f193fcd9170d06d0c77187a2ee18c1` | +| B300 通用对齐 | 恢复设备判断;UE8M0 只在 SM100 上启用,不再无条件应用到所有设备 | `f5f3ed2cc74857ef8821840d5ca0d42c4f2a3e67` | +| DSV4 DP 结束 barrier review | 不纳入无条件 DSV4 barrier;保留此前 CPU-cache 场景已有的条件 barrier | `a2a7052d1a9ffdf765e81f7c43bf59807d484f08` | + +清理后,`lightllm/server/httpserver_for_pd_master/`、`lightllm/server/router/req_queue/`、`lightllm/server/multi_level_kv_cache/disk_cache_worker.py` 和 `lightllm/utils/device_utils.py` 相对拆分基线无额外差异;`communication_op.py` 只保留 DSV4 `experts_` 字段兼容,HTTP server 只保留 Vision 图像块不可跨 prefill 切分的校验。此前确认的 MTP CUDA Graph hidden 输入修复继续保留。 diff --git a/test/benchmark/static_inference/static_benchmark.py b/test/benchmark/static_inference/static_benchmark.py index 4f68933be5..279840d5af 100644 --- a/test/benchmark/static_inference/static_benchmark.py +++ b/test/benchmark/static_inference/static_benchmark.py @@ -6,7 +6,6 @@ import argparse import copy -import json import math import os import queue @@ -29,8 +28,6 @@ sys.path.append(str(REPO_ROOT)) from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput -from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4MemoryManager -from lightllm.common.req_manager import DeepseekV4ReqManager from lightllm.common.basemodel.triton_kernel.gen_mtp_prefill_params import gen_mtp_new_input_ids from lightllm.models import get_draft_model_class, get_model from lightllm.server.api_cli import make_argument_parser @@ -305,7 +302,6 @@ def _fill_mtp_prefill_kv( b_next_token_ids=current_next_ids, mtp_draft_input_hiddens=draft_output.mtp_collector.spec_hidden, ) - draft_input.mtp_draft_input_hiddens = draft_output.mtp_collector.spec_hidden draft_output = self.draft_models[draft_index].forward(draft_input) current_next_ids = self._argmax_ids(draft_output.logits).cuda(non_blocking=True) mtp_candidates.append(current_next_ids.detach().cpu()) @@ -432,9 +428,7 @@ def _run_mtp_draft_decode( model_output: ModelOutput, step_width: int, ): - draft_input = copy.copy(model_input) - draft_input.b_seq_len = model_input.b_seq_len.clone() - draft_input.b_seq_len_cpu = model_input.b_seq_len_cpu.clone() + draft_input = model_input draft_output = model_output draft_next_ids = self._argmax_ids(model_output.logits).cuda(non_blocking=True) generated = [draft_next_ids.detach()] @@ -449,7 +443,6 @@ def _run_mtp_draft_decode( if self.args.mtp_mode.startswith("eagle") and step + 1 < self.args.mtp_step: draft_input.b_seq_len += 1 - draft_input.b_seq_len_cpu += 1 draft_input.max_kv_seq_len += 1 return torch.stack(generated[:step_width], dim=1) @@ -508,34 +501,6 @@ def _materialize_cached_prefix(self, req_idx: torch.Tensor, cached_len: int): if cached_len <= 0: return self._ensure_req_kv_capacity(req_idx, cpu_i32_full((int(req_idx.shape[0]),), cached_len)) - req_idx_gpu = req_idx.cuda(non_blocking=True) - mem_indexes_gpu = self.model.req_manager.req_to_token_indexs[req_idx_gpu, :cached_len] - self._materialize_cached_prefix_extra_slots(req_idx, cached_len, mem_indexes_gpu) - - def _materialize_cached_prefix_extra_slots( - self, req_idx: torch.Tensor, cached_len: int, mem_indexes_gpu: torch.Tensor - ): - req_manager = self.model.req_manager - if not isinstance(req_manager, DeepseekV4ReqManager): - return - batch_size = int(req_idx.shape[0]) - req_list = req_idx.tolist() - seq_list = [cached_len] * batch_size - - swa_ready_len = self._cached_prefix_swa_ready_len(cached_len) - req_manager.prepare_prefill_swa( - req_list=req_list, - ready_list=[swa_ready_len] * batch_size, - seq_list=seq_list, - mem_indexes=mem_indexes_gpu[:, swa_ready_len:].contiguous(), - ) - - def _cached_prefix_swa_ready_len(self, cached_len: int) -> int: - req_manager: DeepseekV4ReqManager = self.model.req_manager - retain_len = int(req_manager._swa_retain_len()) - ready_len = max(0, int(cached_len) - retain_len) - page_size = int(req_manager.get_prompt_cache_page_size()) - return ready_len // page_size * page_size def _make_prefill_input(self, token_chunk: np.ndarray, req_idx: torch.Tensor, ready_cache_len: int) -> ModelInput: batch_size, q_len = token_chunk.shape @@ -1009,38 +974,6 @@ def decode_profile_batch_divisor(args: SimpleNamespace, case: BenchmarkCase) -> return align_up(logical_kv_len, args.page_size) -def filter_capacity_decode_cases( - args: SimpleNamespace, - cases: Sequence[BenchmarkCase], - mem_manager, -) -> List[BenchmarkCase]: - if not getattr(args, "decode_filter_capacity", False): - return list(cases) - - resolved: List[BenchmarkCase] = [] - capacity_tokens = int(mem_manager.size) - for case in cases: - if case.stage != "decode": - resolved.append(case) - continue - divisor = decode_profile_batch_divisor(args, case) - fits_capacity = case.batch_size * divisor <= capacity_tokens - if fits_capacity and isinstance(mem_manager, DeepseekV4MemoryManager) and mem_manager.n_c4 > 0: - c4_page_size = int(mem_manager.c4_pool.page_size) - c4_entries_per_req = divisor // 4 - c4_pages_per_req = (c4_entries_per_req + c4_page_size - 1) // c4_page_size - fits_capacity = case.batch_size * c4_pages_per_req <= mem_manager.c4_num_pages - if fits_capacity: - resolved.append( - replace( - case, - profiled_max_total_token_num=capacity_tokens, - profiled_batch_divisor=divisor, - ) - ) - return resolved - - def resolve_profile_decode_cases( args: SimpleNamespace, cases: Sequence[BenchmarkCase], @@ -1169,12 +1102,7 @@ def normalize_args(args: argparse.Namespace, cases: Sequence[BenchmarkCase]) -> and args.max_total_token_num is None ) prefill_batch_size_needs_profile = args.benchmark in {"all", "prefill"} and args.max_total_token_num is None - decode_capacity_needs_profile = ( - args.benchmark in {"all", "decode"} and args.decode_filter_capacity and args.max_total_token_num is None - ) - needs_profiled_batch_size = ( - decode_batch_size_needs_profile or prefill_batch_size_needs_profile or decode_capacity_needs_profile - ) + needs_profiled_batch_size = decode_batch_size_needs_profile or prefill_batch_size_needs_profile if args.max_total_token_num is None and not needs_profiled_batch_size: tokens_per_req = align_up(args.max_req_total_len + mtp_width + 8, args.page_size) @@ -1240,7 +1168,7 @@ def build_model_kvargs(args: SimpleNamespace, rank_id: int) -> Dict: "max_req_num": max(args.running_max_req_size, args.graph_max_batch_size), "batch_max_tokens": args.batch_max_tokens, "run_mode": "normal", - "max_seq_length": args.max_req_total_len + max(8, args.mtp_step * 2), + "max_seq_length": args.max_req_total_len, "disable_cudagraph": args.disable_cudagraph, "llm_prefill_att_backend": args.llm_prefill_att_backend, "llm_decode_att_backend": args.llm_decode_att_backend, @@ -1345,9 +1273,6 @@ def run_worker(args_dict: Dict, case_dicts: List[Dict], rank_id: int, ans_queue) model, _ = get_model(model_cfg, model_kvargs) cases = resolve_batch_max_prefill_cases(args, cases, model.mem_manager.size) cases = resolve_profile_decode_cases(args, cases, model.mem_manager.size) - cases = filter_capacity_decode_cases(args, cases, model.mem_manager) - if not cases: - raise ValueError("no benchmark cases remain after capacity filtering") if defer_cudagraph: init_deferred_cudagraph(args, cases, model_kvargs, model) draft_models = init_mtp_draft_models(args, model_kvargs, model) @@ -1659,11 +1584,6 @@ def add_static_benchmark_args(parser: argparse.ArgumentParser): "from profiled max_total_token_num per context" ), ) - parser.add_argument( - "--decode_filter_capacity", - action="store_true", - help="drop explicit decode cases whose batch size cannot fit profiled KV capacity", - ) parser.add_argument( "--mtp_accept_rate", type=float, @@ -1673,7 +1593,6 @@ def add_static_benchmark_args(parser: argparse.ArgumentParser): parser.add_argument("--warmup_iters", type=int, default=1) parser.add_argument("--bench_iters", type=int, default=1) parser.add_argument("--seed", type=int, default=1234) - parser.add_argument("--dump_file", type=str, default=None, help="write aggregated benchmark results as JSON") def main(argv: Optional[Sequence[str]] = None): @@ -1689,12 +1608,7 @@ def main(argv: Optional[Sequence[str]] = None): args = normalize_args(args, cases) set_env_start_args(args) - results = run_benchmark(args, cases) - if args.dump_file and args.node_rank == 0: - dump_path = Path(args.dump_file) - dump_path.parent.mkdir(parents=True, exist_ok=True) - payload = {"args": vars(args), "results": results} - dump_path.write_text(json.dumps(payload, indent=2, sort_keys=True)) + run_benchmark(args, cases) if __name__ == "__main__": diff --git a/unit_tests/common/test_deepseek4_paged_cache.py b/unit_tests/common/test_deepseek4_paged_cache.py index 176f357f59..9823eff337 100644 --- a/unit_tests/common/test_deepseek4_paged_cache.py +++ b/unit_tests/common/test_deepseek4_paged_cache.py @@ -45,7 +45,7 @@ def cache(monkeypatch, tmp_path): compress_rates=[4, 128, 0], max_request_num=2, mtp_step=3, - swa_full_tokens_ratio=1.0, + swa_page_num=64, ) requests = DeepseekV4ReqManager(2, 4096, manager, sliding_window=128) yield manager, requests @@ -628,7 +628,7 @@ def _checkpoint_transfer_worker(rank, rendezvous): torch.cuda.set_device(rank) dist.init_process_group( - "nccl", init_method="file://" + rendezvous, rank=rank, world_size=2, timeout=timedelta(seconds=60) + "gloo", init_method="file://" + rendezvous, rank=rank, world_size=2, timeout=timedelta(seconds=60) ) try: buffers = DeepseekV4StateCacheManager(2, DeepseekV4CpuCacheLayout.from_compress_rates([4, 128])) @@ -638,7 +638,7 @@ def _checkpoint_transfer_worker(rank, rendezvous): module = DPKVSharedMoudle.__new__(DPKVSharedMoudle) module.backend = SimpleNamespace( args=SimpleNamespace(linear_att_hash_page_size=256, linear_att_page_block_num=8, max_req_total_len=8192), - node_nccl_group=dist.group.WORLD, + node_gloo_group=dist.group.WORLD, model=SimpleNamespace(mem_manager=SimpleNamespace(big_page_buffers=buffers)), radix_cache=SimpleNamespace(get_big_page_ids_by_node=lambda node: [0, 1]), ) @@ -671,6 +671,6 @@ def _checkpoint_transfer_worker(rank, rendezvous): dist.destroy_process_group() -@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires two CUDA devices for NCCL") +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires two CUDA devices") def test_dp_checkpoint_transfer_between_processes(tmp_path): torch.multiprocessing.spawn(_checkpoint_transfer_worker, args=(str(tmp_path / "rendezvous"),), nprocs=2, join=True) diff --git a/unit_tests/server/core/objs/test_req.py b/unit_tests/server/core/objs/test_req.py index d3bad21df5..3cae97ac51 100644 --- a/unit_tests/server/core/objs/test_req.py +++ b/unit_tests/server/core/objs/test_req.py @@ -20,7 +20,6 @@ def setup_module_env(): "llm_decode_att_backend": ["None"], "cpu_cache_token_page_size": 256, "enable_cpu_cache": False, - "enable_ep_moe": False, "model_dir": "", "page_size": 4, } diff --git a/unit_tests/server/httpserver/test_pd_compact_transport.py b/unit_tests/server/httpserver/test_pd_compact_transport.py deleted file mode 100644 index 984d55a65b..0000000000 --- a/unit_tests/server/httpserver/test_pd_compact_transport.py +++ /dev/null @@ -1,73 +0,0 @@ -import asyncio -from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock - -from lightllm.server.httpserver.pd_loop import _pd_process_generate -from lightllm.server.pd_io_struct import ( - PD_COMPACT_TOKEN_INFO_LEN, - build_pd_compact_token_info, - unpack_pd_compact_token_info, -) - - -def _compact_metadata(token_id): - return { - "id": token_id, - "count_output_tokens": 1, - "prompt_tokens": 4, - "prompt_cache_len": 2, - "mtp_accepted_token_num": 0, - "mtp_verify_token_num": 0, - "mtp_verify_step_num": 0, - } - - -def test_pd_compact_token_roundtrip_preserves_token_id(): - finish_status = SimpleNamespace(status=1) - - for token_id in (42, None): - metadata = _compact_metadata(token_id) - compact = build_pd_compact_token_info(123, "token", metadata, finish_status) - - assert len(compact) == PD_COMPACT_TOKEN_INFO_LEN - assert unpack_pd_compact_token_info(compact) == (123, "token", metadata, finish_status.status) - - -def test_pd_transport_without_logprobs_keeps_token_id_in_compact_packet(): - metadata = { - **_compact_metadata(42), - "logprob": -0.1, - "cumlogprob": -0.1, - "special": False, - "logprobs": {42: -0.1}, - } - finish_status = SimpleNamespace(status=1) - - class _Manager: - args = SimpleNamespace(run_mode="decode") - - async def generate(self, **_kwargs): - yield 123, "token", metadata, finish_status - - async def run(): - forwarding_queue = AsyncMock() - await _pd_process_generate( - manager=_Manager(), - prompt="prompt", - sampling_params=SimpleNamespace(group_request_id=123, return_output_logprobs=False), - multimodal_params=MagicMock(), - forwarding_queue=forwarding_queue, - pd_upload_websocket=AsyncMock(), - pd_event=asyncio.Event(), - ) - return forwarding_queue.put.await_args.args[0] - - compact = asyncio.run(run()) - assert len(compact) == PD_COMPACT_TOKEN_INFO_LEN - - _, _, unpacked_metadata, _ = unpack_pd_compact_token_info(compact) - assert unpacked_metadata["id"] == 42 - assert "logprob" not in unpacked_metadata - assert "cumlogprob" not in unpacked_metadata - assert "special" not in unpacked_metadata - assert "logprobs" not in unpacked_metadata diff --git a/unit_tests/server/httpserver/test_pd_generate_error.py b/unit_tests/server/httpserver/test_pd_generate_error.py index 26c93d81e0..ca4515ca15 100644 --- a/unit_tests/server/httpserver/test_pd_generate_error.py +++ b/unit_tests/server/httpserver/test_pd_generate_error.py @@ -36,7 +36,7 @@ class _SuccessfulManager: args = SimpleNamespace(run_mode="prefill") async def generate(self, **_kwargs): - yield 123, "token", {"count_output_tokens": 1}, FinishStatus(FinishStatus.FINISHED_STOP) + yield 123, "token", {}, FinishStatus(FinishStatus.FINISHED_STOP) class _StopPrefillManager: diff --git a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py index 9908b9d49c..bdb877773e 100644 --- a/unit_tests/server/router/dynamic_prompt/test_radix_cache.py +++ b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py @@ -249,12 +249,7 @@ def test_case10(): 测试场景:测试 flush_cache 函数 """ print("\nTest Case 10: Testing flush_cache function\n") - allocator = SimpleNamespace(can_use_mem_size=95) - - def free(indexes): - allocator.can_use_mem_size += len(indexes) - - tree = RadixCache(100, 0, mem_manager=SimpleNamespace(page_size=1, allocator=allocator, free=free)) + tree = RadixCache(100, 0) tree.insert(torch.tensor([1, 2, 3], dtype=torch.int64)) tree.insert(torch.tensor([1, 2, 3, 4, 5], dtype=torch.int64)) tree_node, size, values = tree.match_prefix( @@ -262,9 +257,7 @@ def free(indexes): ) assert tree_node is not None assert size == 3 - tree.dec_node_ref_counter(tree_node) tree.flush_cache() - assert allocator.can_use_mem_size == 100 tree_node, size, values = tree.match_prefix( torch.tensor([1, 2, 3], dtype=torch.int64, device="cpu"), update_refs=True ) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py deleted file mode 100644 index 51358bd1f8..0000000000 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_control.py +++ /dev/null @@ -1,77 +0,0 @@ -from types import SimpleNamespace - -import torch - -from lightllm.server.router.model_infer.mode_backend import base_backend as base_backend_module -from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend -from lightllm.server.router.model_infer.mode_backend.dp_backend.control_state import DPControlState, RunWay - - -def test_dp_req_presence_uses_one_cpu_collective(monkeypatch): - control_group = object() - backend = ModeBackend.__new__(ModeBackend) - backend.dp_control_tensor = torch.zeros(2, dtype=torch.int32, device="cpu") - calls = [] - - def all_reduce(tensor, op, group, async_op): - calls.append((tensor.tolist(), tensor.device.type, op, group, async_op)) - tensor.fill_(1) - - monkeypatch.setattr(base_backend_module.dist_group_manager, "dp_control_group", control_group) - monkeypatch.setattr(base_backend_module, "all_reduce", all_reduce) - - has_prefill, has_decode = backend._dp_all_reduce_req_presence(prefill_reqs=[object()], decode_reqs=[]) - - assert (has_prefill, has_decode) == (True, True) - assert calls == [ - ([1, 0], "cpu", torch.distributed.ReduceOp.MAX, control_group, False), - ] - - -def test_new_request_readiness_uses_cpu_control_group(monkeypatch): - control_group = object() - backend = ModeBackend.__new__(ModeBackend) - backend.is_master_in_node = True - backend.shm_reqs_io_buffer = SimpleNamespace(is_ready=lambda: True) - backend.node_broadcast_tensor = torch.zeros(1, dtype=torch.int32, device="cpu") - backend.node_gloo_group = control_group - backend.node_world_size = 8 - backend.args = SimpleNamespace(node_rank=0) - backend.is_pd_mode = False - calls = [] - - def broadcast(tensor, src, group, async_op): - calls.append((tensor.tolist(), tensor.device.type, src, group, async_op)) - - backend._read_reqs_buffer_and_init_reqs = lambda: calls.append("read") - monkeypatch.setattr(base_backend_module, "broadcast", broadcast) - - backend._try_read_new_reqs_normal() - - assert calls == [ - ([1], "cpu", 0, control_group, False), - "read", - ] - - -def test_aggressive_dp_control_prefers_prefill(): - state = DPControlState.__new__(DPControlState) - state.is_aggressive_schedule = True - state.step_count = 0 - - assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.PREFILL - assert state.select_run_way(has_prefill=False, has_decode=True) is RunWay.DECODE - assert state.select_run_way(has_prefill=False, has_decode=False) is RunWay.PASS - assert state.step_count == 3 - - -def test_normal_dp_control_preserves_decode_budget(): - state = DPControlState.__new__(DPControlState) - state.is_aggressive_schedule = False - state.decode_max_step = 1 - state.left_decode_num = 1 - state.step_count = 0 - - assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.DECODE - assert state.select_run_way(has_prefill=True, has_decode=True) is RunWay.PREFILL - assert state.left_decode_num == 1 diff --git a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py index 39fb0fc501..af9dfa9534 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_dp_overlap_spec_engine.py @@ -309,6 +309,8 @@ def test_dp_overlap_engine_delegates_raw_verify_layout_to_proposer(): calls = {} class _Proposer: + backend = SimpleNamespace(is_deepseek_v4=False) + def propose_next_overlap(self, **kwargs): calls.update(kwargs) return SpecProposal(token_ids=kwargs["target_next_token_ids0"].new_empty((3, 7))) diff --git a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py index 09de7cbbd5..260dbe1065 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py @@ -122,9 +122,11 @@ def start_offload(req, cpu_kv_cache_stream): ) for index, (length, req_cache_tiers) in enumerate(zip((100, 200, 300), cache_tiers)) ] - true_finished_reqs = module.offload_finished_reqs_to_cpu_cache(reqs) + offload_reqs = [req for req in reqs if CacheTier.CPU in req.cache_tiers or CacheTier.DISK in req.cache_tiers] - assert true_finished_reqs == [reqs[0]] + true_finished_reqs = module.offload_finished_reqs_to_cpu_cache(offload_reqs) + + assert true_finished_reqs == [] assert offload_calls == [(1, False), (2, True)] assert len(module.cpu_cache_handle_queue) == 2 diff --git a/unit_tests/server/test_pd_start_args.py b/unit_tests/server/test_pd_start_args.py index deadbd8fe0..3e98fa56a4 100644 --- a/unit_tests/server/test_pd_start_args.py +++ b/unit_tests/server/test_pd_start_args.py @@ -19,6 +19,8 @@ def test_dsv4_page_and_checkpoint_configuration(monkeypatch, page_size, hash_pag monkeypatch.setattr("lightllm.server.api_start.is_hybrid_att_model", lambda model_dir: True) def validation_finished(args): + assert args.page_size == 256 + assert args.linear_att_hash_page_size == 256 assert args.cpu_cache_token_page_size == 2048 raise RuntimeError("DSV4 page-size validation passed") @@ -36,10 +38,7 @@ def validation_finished(args): disable_audio=True, disable_shm_warning=True, ) - if page_size != 256 or hash_page_size != 256: - with pytest.raises(ValueError, match="DeepSeek-V4 requires"): - _launch_subprocesses(args) - elif cpu_page_size is not None: + if cpu_page_size is not None: with pytest.raises(ValueError, match="CPU cache pages must match"): _launch_subprocesses(args) else: From 3f41c74c10634cedc2a8445a5844a445f5184256 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 08:16:56 +0000 Subject: [PATCH 206/214] minor --- lightllm/server/httpserver/manager.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/lightllm/server/httpserver/manager.py b/lightllm/server/httpserver/manager.py index 1eeafea37e..420c55eece 100644 --- a/lightllm/server/httpserver/manager.py +++ b/lightllm/server/httpserver/manager.py @@ -375,14 +375,6 @@ async def generate( await self._log_req_header(request_headers, group_request_id) # encode prompt_ids = await self._encode(prompt, multimodal_params, sampling_params) - for image in multimodal_params.images: - if image.block_start_idx is not None: - block_token_num = image.block_end_idx - image.block_start_idx - if block_token_num > self.args.batch_max_tokens: - raise ValueError( - f"image prefill block token count {block_token_num} exceeds " - f"batch_max_tokens={self.args.batch_max_tokens}; increase --batch_max_tokens" - ) self._log_stage_timing( group_request_id, start_time, From e6877850a78cd1fa76b3e5c6670fd9400d89d578 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 16:35:08 +0800 Subject: [PATCH 207/214] refactor(moe): unify clamped SwiGLU activation args --- .../fused_moe/fused_moe_weight.py | 24 ++++++++++++---- .../fused_moe/impl/deepgemm_impl.py | 24 ++++++++++++---- .../fused_moe/impl/triton_impl.py | 8 ++++-- .../fused_moe/grouped_fused_moe_ep.py | 2 +- .../fused_moe/moe_silu_and_mul.py | 11 +------- .../moe_silu_and_mul_mix_quant_ep.py | 22 ++++++--------- .../layer_infer/transformer_layer_infer.py | 28 +++++++++++++++---- .../fused_moe/test_activation_config.py | 7 +++++ 8 files changed, 82 insertions(+), 44 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 4dab4eab98..9a37767265 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -171,7 +171,9 @@ def experts_with_topk( topk_ids: torch.Tensor, is_prefill: Optional[bool] = None, infer_state=None, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, alloc_tensor_func=torch.empty, ) -> torch.Tensor: moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) @@ -184,7 +186,9 @@ def experts_with_topk( topk_weights=topk_weights, topk_ids=topk_ids, is_prefill=is_prefill, - clamp_limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, alloc_tensor_func=alloc_tensor_func, ) @@ -263,7 +267,9 @@ def masked_group_gemm( masked_m: torch.Tensor, dtype: torch.dtype, expected_m: int, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, ): assert self.enable_ep_moe, "masked_group_gemm is only supported when enable_ep_moe is True" return self.fuse_moe_impl.masked_group_gemm( @@ -273,7 +279,9 @@ def masked_group_gemm( masked_m=masked_m, dtype=dtype, expected_m=expected_m, - clamp_limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, ) def prefilled_group_gemm( @@ -286,7 +294,9 @@ def prefilled_group_gemm( recv_topk_weights: torch.Tensor, hidden_dtype=torch.bfloat16, microbatch_index: int = 0, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, ): assert self.enable_ep_moe, "prefilled_group_gemm is only supported when enable_ep_moe is True" return self.fuse_moe_impl.prefilled_group_gemm( @@ -300,7 +310,9 @@ def prefilled_group_gemm( w2=self.w2, hidden_dtype=hidden_dtype, microbatch_index=microbatch_index, - clamp_limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, ) def low_latency_combine( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index cb6b6e65b7..2598cdd61d 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -217,7 +217,9 @@ def masked_group_gemm( masked_m: torch.Tensor, dtype: torch.dtype, expected_m: int, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -230,7 +232,9 @@ def masked_group_gemm( w2_weight, w2_scale, expected_m=expected_m, - limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, ) def prefilled_group_gemm( @@ -245,7 +249,9 @@ def prefilled_group_gemm( w2: WeightPack, hidden_dtype=torch.bfloat16, microbatch_index: int = 0, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -265,7 +271,9 @@ def prefilled_group_gemm( block_size_k=self.quant_method.block_size, workspace=dist_group_manager.get_deep_ep_prefill_moe_workspace(microbatch_index), hidden_dtype=hidden_dtype, - limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, ) else: gather_out = torch.empty( @@ -281,7 +289,13 @@ def prefilled_group_gemm( N = w13_weight.shape[1] _gemm_out_a = torch.zeros((1, N), device=recv_x[0].device, dtype=hidden_dtype) _silu_out = torch.zeros((1, N // 2), device=recv_x[0].device, dtype=hidden_dtype) - silu_and_mul_fwd(_gemm_out_a.view(-1, N), _silu_out, limit=clamp_limit) + silu_and_mul_fwd( + _gemm_out_a.view(-1, N), + _silu_out, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, + ) _gemm_out_a, _silu_out = None, None del recv_x return gather_out diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index adbb46a7dd..83500d4eae 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -99,7 +99,9 @@ def fused_experts_with_topk( topk_weights: torch.Tensor, topk_ids: torch.Tensor, is_prefill: Optional[bool] = None, - clamp_limit: Optional[float] = None, + alpha: Optional[float] = None, + limit: Optional[float] = None, + clamp_up_add_one: bool = True, alloc_tensor_func=torch.empty, ): return self._fused_experts( @@ -109,7 +111,9 @@ def fused_experts_with_topk( topk_weights=topk_weights, topk_ids=topk_ids, is_prefill=is_prefill, - limit=clamp_limit, + alpha=alpha, + limit=limit, + clamp_up_add_one=clamp_up_add_one, alloc_tensor_func=alloc_tensor_func, ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index 8f86abe915..b26e8fa668 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -221,7 +221,7 @@ def fused_experts( clamp_up_add_one: bool = True, alloc_tensor_func: Callable = torch.empty, ): - assert alpha is None or limit is not None + assert (limit is None and alpha is None) or (limit is not None and alpha is not None) check_ep_expert_dtype(quant_method) if use_sm100_mega_moe(quant_method): if limit is not None: diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py index 1de69f79a4..8169cefc3c 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul.py @@ -24,7 +24,6 @@ def _silu_and_mul_kernel_fast( NEED_MASK: tl.constexpr, layout: tl.constexpr = "blocked", # "blocked" or "interleaved" USE_LIMIT_AND_ALPHA: tl.constexpr = False, - USE_LIMIT_ONLY: tl.constexpr = False, CLAMP_UP_ADD_ONE: tl.constexpr = True, USE_TANH_APPROXIMATE_GELU: tl.constexpr = False, ): @@ -76,11 +75,6 @@ def _silu_and_mul_kernel_fast( up += 1 tl.store(output_ptr + out_offsets, up * gate, mask=mask) else: - if USE_LIMIT_ONLY: - # clamped swiglu (DeepSeek-V4 swiglu_limit): clamp 后接标准 silu, - # 无 gpt-oss 的 alpha 缩放与 (up+1)。 - gate = tl.minimum(gate, limit) - up = tl.minimum(tl.maximum(up, -limit), limit) if USE_TANH_APPROXIMATE_GELU: # tanh-approx GELU, matching Gemma's gelu_pytorch_tanh MLP. gate_cubed = gate * gate * gate @@ -130,8 +124,7 @@ def silu_and_mul_fwd( ): assert input.stride(-1) == 1 assert output.is_contiguous() - # limit+alpha: gpt-oss 语义 (up+1)*silu(alpha*gate); 仅 limit: clamp 后标准 silu (DeepSeek-V4) - assert alpha is None or limit is not None + assert (limit is None and alpha is None) or (limit is not None and alpha is not None) stride_input_m = input.stride(0) stride_input_n = input.stride(1) @@ -154,7 +147,6 @@ def silu_and_mul_fwd( while triton.cdiv(size_m, BLOCK_M) > 8192: BLOCK_M *= 2 USE_LIMIT_AND_ALPHA = limit is not None and alpha is not None - USE_LIMIT_ONLY = limit is not None and alpha is None grid = ( triton.cdiv(size_n, BLOCK_N), @@ -179,7 +171,6 @@ def silu_and_mul_fwd( num_warps=num_warps, layout=layout, USE_LIMIT_AND_ALPHA=USE_LIMIT_AND_ALPHA, - USE_LIMIT_ONLY=USE_LIMIT_ONLY, CLAMP_UP_ADD_ONE=clamp_up_add_one, USE_TANH_APPROXIMATE_GELU=ffn_use_tanh_approximate_gelu(), ) diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py index e7a89abb2e..661b267088 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/moe_silu_and_mul_mix_quant_ep.py @@ -26,7 +26,6 @@ def _silu_and_mul_post_quant_kernel( fp8_min, BLOCK_N: tl.constexpr, NUM_STAGE: tl.constexpr, - USE_LIMIT_ONLY: tl.constexpr = False, USE_TANH_APPROXIMATE_GELU: tl.constexpr = False, USE_LIMIT_AND_ALPHA: tl.constexpr = False, alpha: tl.constexpr = None, @@ -62,17 +61,13 @@ def _silu_and_mul_post_quant_kernel( gate = gate / (1 + tl.exp(-gate * alpha)) if CLAMP_UP_ADD_ONE: up += 1 + elif USE_TANH_APPROXIMATE_GELU: + gate_cubed = gate * gate * gate + tanh_arg = 0.7978845608028654 * (gate + 0.044715 * gate_cubed) + tanh_val = 2.0 / (1.0 + tl.exp(-2.0 * tanh_arg)) - 1.0 + gate = 0.5 * gate * (1.0 + tanh_val) else: - if USE_LIMIT_ONLY: - gate = tl.minimum(gate, limit) - up = tl.minimum(tl.maximum(up, -limit), limit) - if USE_TANH_APPROXIMATE_GELU: - gate_cubed = gate * gate * gate - tanh_arg = 0.7978845608028654 * (gate + 0.044715 * gate_cubed) - tanh_val = 2.0 / (1.0 + tl.exp(-2.0 * tanh_arg)) - 1.0 - gate = 0.5 * gate * (1.0 + tanh_val) - else: - gate = gate / (1 + tl.exp(-gate)) + gate = gate / (1 + tl.exp(-gate)) gate = gate.to(input_ptr.dtype.element_ty) gate_up = up * gate if USE_LIMIT_AND_ALPHA: @@ -110,7 +105,7 @@ def silu_and_mul_masked_post_quant_fwd( masked_m shape [expert_num], """ - assert alpha is None or limit is not None + assert (limit is None and alpha is None) or (limit is not None and alpha is not None) assert input.is_contiguous() assert output.dtype == torch.float8_e4m3fn assert output.is_contiguous() @@ -157,9 +152,8 @@ def silu_and_mul_masked_post_quant_fwd( fp8_min, BLOCK_N=BLOCK_N, NUM_STAGE=NUM_STAGES, - USE_LIMIT_ONLY=limit is not None and alpha is None, USE_TANH_APPROXIMATE_GELU=ffn_use_tanh_approximate_gelu(), - USE_LIMIT_AND_ALPHA=limit is not None and alpha is not None, + USE_LIMIT_AND_ALPHA=limit is not None, alpha=alpha, limit=limit, CLAMP_UP_ADD_ONE=clamp_up_add_one, diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index c51ccb0e1c..9e3c92d849 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -203,7 +203,9 @@ def overlap_tpsp_context_forward( recv_weights0, hidden_dtype=x0.dtype, microbatch_index=infer_state.microbatch_index, - clamp_limit=self.swiglu_limit, + alpha=1.0, + limit=self.swiglu_limit, + clamp_up_add_one=False, ) combine_event0 = ElasticBuffer.capture() @@ -219,7 +221,9 @@ def overlap_tpsp_context_forward( recv_weights1, hidden_dtype=x1.dtype, microbatch_index=infer_state1.microbatch_index, - clamp_limit=self.swiglu_limit, + alpha=1.0, + limit=self.swiglu_limit, + clamp_up_add_one=False, ) combine_event1 = ElasticBuffer.capture() @@ -294,7 +298,9 @@ def overlap_tpsp_token_forward( masked_m0, x0.dtype, expected_m, - clamp_limit=self.swiglu_limit, + alpha=1.0, + limit=self.swiglu_limit, + clamp_up_add_one=False, ) dispatch_hook1() @@ -305,7 +311,9 @@ def overlap_tpsp_token_forward( masked_m1, x1.dtype, expected_m, - clamp_limit=self.swiglu_limit, + alpha=1.0, + limit=self.swiglu_limit, + clamp_up_add_one=False, ) combine_hook0() @@ -540,7 +548,9 @@ def _routed_experts( topk_ids=indices, is_prefill=infer_state.is_prefill, infer_state=infer_state, - clamp_limit=float(self.swiglu_limit), + alpha=1.0, + limit=float(self.swiglu_limit), + clamp_up_add_one=False, alloc_tensor_func=self.alloc_tensor, ) @@ -548,7 +558,13 @@ def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: De input = input.view(-1, self.embed_dim_) gate_up = layer_weight.gate_up_proj.mm(input) shared = self.alloc_tensor((input.size(0), gate_up.size(1) // 2), input.dtype) - silu_and_mul_fwd(gate_up, shared, limit=self.swiglu_limit) + silu_and_mul_fwd( + gate_up, + shared, + alpha=1.0, + limit=self.swiglu_limit, + clamp_up_add_one=False, + ) input = None gate_up = None out = layer_weight.down_proj.mm(shared) diff --git a/unit_tests/common/fused_moe/test_activation_config.py b/unit_tests/common/fused_moe/test_activation_config.py index 311b1fc652..8972dd88f0 100644 --- a/unit_tests/common/fused_moe/test_activation_config.py +++ b/unit_tests/common/fused_moe/test_activation_config.py @@ -87,7 +87,14 @@ def test_call_parameters_reach_expert_activation(monkeypatch, activation): if kwargs: default_output = weight.experts(x.clone(), router, 2, True, False, 0, 0) actual = weight.experts(x.clone(), router, 2, True, False, 0, 0, **kwargs) + actual_with_topk = weight.experts_with_topk( + x.clone(), + topk_weights=probs, + topk_ids=top.indices, + **kwargs, + ) torch.testing.assert_close(actual, expected.bfloat16(), atol=0.125, rtol=0.01) + torch.testing.assert_close(actual_with_topk, expected.bfloat16(), atol=0.125, rtol=0.01) if kwargs: # A clamped call must not change subsequent calls on the same weight. actual_default = weight.experts(x.clone(), router, 2, True, False, 0, 0) From 252c78461d327d4204944ef0f8df26cdedf79c4b Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Mon, 28 Sep 2026 16:35:51 +0800 Subject: [PATCH 208/214] Revert "warmup tilelang" This reverts commit fc8b7ec411774fab269e1e3799efff4ac15f826e. Preserve the hc_post import required by the subsequent MTP hidden preparation path. Record the rollback in revert.md. --- lightllm/common/basemodel/basemodel.py | 5 -- .../layer_infer/hyper_connection.py | 14 ---- lightllm/models/deepseek_v4/model.py | 65 +------------------ revert.md | 1 + 4 files changed, 2 insertions(+), 83 deletions(-) diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py index 51e4c48a3f..30865f38cd 100755 --- a/lightllm/common/basemodel/basemodel.py +++ b/lightllm/common/basemodel/basemodel.py @@ -137,7 +137,6 @@ def __init__(self, kvargs): self._init_hidden_collector() self._autotune_warmup() - self._kernel_warmup() self._init_padded_req() self._init_cudagraph() self._init_prefill_cuda_graph() @@ -318,10 +317,6 @@ def _init_prefill_cuda_graph(self): def _init_custom(self): pass - def _kernel_warmup(self): - """Warm model-specific kernels before CUDA graph capture.""" - return - def _init_hidden_collector(self): self.hidden_collector_prototype = self.mtp_manager.create_hidden_collector(model=self) diff --git a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py index 920309f53d..c2592c4d0b 100644 --- a/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py +++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py @@ -1,5 +1,4 @@ import torch -from vllm._tilelang_ops import compute_num_split try: import vllm.model_executor.layers.mhc # noqa: F401 @@ -11,19 +10,6 @@ HC_POST_ALPHA = 2.0 -def mhc_warmup_token_sizes(max_tokens, hidden_size, hc_mult): - """Choose one token count for every reachable mHC split-K variant.""" - block_size = 64 - token_sizes_by_split = {} - - max_grid_size = (max_tokens + block_size - 1) // block_size - for grid_size in range(1, max_grid_size + 1): - n_splits = compute_num_split(block_size, hc_mult * hidden_size, grid_size) - token_sizes_by_split.setdefault(n_splits, min(grid_size * block_size, max_tokens)) - - return sorted(token_sizes_by_split.values()) - - def hc_pre(residual, hc_fn, hc_scale, hc_base, rms_eps, hc_eps, sinkhorn_iters, norm_weight, norm_eps): """Standalone hc_pre for the first layer. residual:[T, hc, dim] -> (x[T,dim], residual, post_mix[T,hc,1], res_mix[T,hc,hc]); the sub-layer RMSNorm is fused via norm_weight.""" diff --git a/lightllm/models/deepseek_v4/model.py b/lightllm/models/deepseek_v4/model.py index 255f649b69..a90696b360 100644 --- a/lightllm/models/deepseek_v4/model.py +++ b/lightllm/models/deepseek_v4/model.py @@ -3,7 +3,6 @@ import json import math import os -import time import torch from lightllm.models.llama.model import LlamaTpPartModel @@ -34,11 +33,7 @@ from lightllm.common.basemodel.attention.nsa.dsv4_fp8_flashmla_sparse import DSV4_NSA_BACKENDS from lightllm.models.deepseek_v4.infer_struct import DeepseekV4InferStateInfo from lightllm.models.deepseek_v4.workspace import DeepseekV4Workspace -from lightllm.models.deepseek_v4.layer_infer.hyper_connection import ( - hc_head, - hc_post, - mhc_warmup_token_sizes, -) +from lightllm.models.deepseek_v4.layer_infer.hyper_connection import hc_post from lightllm.models.llama.yarn_rotary_utils import ( find_correction_range, linear_ramp_mask, @@ -201,64 +196,6 @@ def _init_custom(self): ) return - @torch.no_grad() - def _kernel_warmup(self): - if self.is_mtp_draft_model: - return - - layer_infer = self.layers_infer[0] - layer_weight = self.trans_layers_weight[0] - hidden_size = self.config["hidden_size"] - hc_mult = self.config["hc_mult"] - split_token_sizes = mhc_warmup_token_sizes( - max_tokens=self.batch_max_tokens, - hidden_size=hidden_size, - hc_mult=hc_mult, - ) - token_sizes = sorted(set(split_token_sizes + [size for size in (1, 8, 17) if size <= self.batch_max_tokens])) - - started = time.perf_counter() - logger.info( - "warming DeepSeek-V4 mHC TileLang kernels for token sizes %s", - token_sizes, - ) - residual = torch.zeros( - max(token_sizes), - hc_mult, - hidden_size, - dtype=torch.bfloat16, - device="cuda", - ) - for token_size in split_token_sizes: - layer_infer._hc_attn_in(residual[:token_size], layer_weight) - - for token_size in (size for size in (1, 8, 17) if size <= self.batch_max_tokens): - hc_state = layer_infer._hc_attn_in(residual[:token_size], layer_weight) - hc_state = layer_infer._hc_ffn_in(*hc_state, layer_weight) - - streams = hc_post(*hc_state) - hc_head( - streams, - self.pre_post_weight.hc_head_fn_.weight, - self.pre_post_weight.hc_head_scale_.weight, - self.pre_post_weight.hc_head_base_.weight, - hc_mult, - hidden_size, - self.config["rms_norm_eps"], - self.config.get("hc_eps", 1e-6), - torch.empty, - ) - - torch.cuda.synchronize() - del residual, hc_state, streams - torch.cuda.empty_cache() - logger.info( - "DeepSeek-V4 mHC TileLang warmup finished in %.2f seconds (%d split-K variants)", - time.perf_counter() - started, - len(split_token_sizes), - ) - return - def prepare_mtp_layer_hidden(self, layer_index: int, hidden): """Materialize the official per-layer DSpark feature from the deferred mHC state.""" if isinstance(hidden, tuple): diff --git a/revert.md b/revert.md index 4039c36036..343be7eb5b 100644 --- a/revert.md +++ b/revert.md @@ -63,5 +63,6 @@ | DeepEP 环境自动调参 | 回退 decode dispatch capacity、NVSHMEM QP depth 等通用自动派生,恢复原有默认值 | `9a071ca5e0e53bab4bca63586f78e92b493cfb68`、`97ae2d12acdd77bcf87a786a710c0f1f049d8196`、`69c260f4a5f193fcd9170d06d0c77187a2ee18c1` | | B300 通用对齐 | 恢复设备判断;UE8M0 只在 SM100 上启用,不再无条件应用到所有设备 | `f5f3ed2cc74857ef8821840d5ca0d42c4f2a3e67` | | DSV4 DP 结束 barrier review | 不纳入无条件 DSV4 barrier;保留此前 CPU-cache 场景已有的条件 barrier | `a2a7052d1a9ffdf765e81f7c43bf59807d484f08` | +| mHC TileLang 启动预热 | 回退通用 `_kernel_warmup` hook、DSV4 mHC 预热流程及 split-K token 枚举;保留 MTP hidden 准备所需的 `hc_post` 引用 | `fc8b7ec411774fab269e1e3799efff4ac15f826e` (`warmup tilelang`) | 清理后,`lightllm/server/httpserver_for_pd_master/`、`lightllm/server/router/req_queue/`、`lightllm/server/multi_level_kv_cache/disk_cache_worker.py` 和 `lightllm/utils/device_utils.py` 相对拆分基线无额外差异;`communication_op.py` 只保留 DSV4 `experts_` 字段兼容,HTTP server 只保留 Vision 图像块不可跨 prefill 切分的校验。此前确认的 MTP CUDA Graph hidden 输入修复继续保留。 From 4cbcd45d4efed5d374232e642fb3fff1307e710c Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 29 Sep 2026 01:14:17 +0000 Subject: [PATCH 209/214] delete alloc_tensor_func --- .../fused_moe/fused_moe_weight.py | 2 -- .../fused_moe/impl/deepgemm_impl.py | 2 -- .../fused_moe/impl/marlin_impl.py | 1 - .../fused_moe/impl/triton_impl.py | 3 --- .../fused_moe/grouped_fused_moe_ep.py | 19 +++++-------------- .../quantization/fp8act_quant_kernel.py | 2 -- .../layer_infer/transformer_layer_infer.py | 1 - .../server/core/objs/py_sampling_params.py | 8 ++++---- lightllm/server/core/objs/sampling_params.py | 8 ++++---- revert.md | 1 + 10 files changed, 14 insertions(+), 33 deletions(-) diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py index 9a37767265..72311d22e4 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/fused_moe_weight.py @@ -174,7 +174,6 @@ def experts_with_topk( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func=torch.empty, ) -> torch.Tensor: moe_capture_callback = get_moe_capture_callback(infer_state, self.layer_num_) if moe_capture_callback is not None: @@ -189,7 +188,6 @@ def experts_with_topk( alpha=alpha, limit=limit, clamp_up_add_one=clamp_up_add_one, - alloc_tensor_func=alloc_tensor_func, ) def low_latency_dispatch( diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py index 2598cdd61d..b8b58d97ed 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/deepgemm_impl.py @@ -79,7 +79,6 @@ def _fused_experts( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func=torch.empty, ): output = fused_experts( hidden_states=input_tensor, @@ -94,7 +93,6 @@ def _fused_experts( alpha=alpha, limit=limit, clamp_up_add_one=clamp_up_add_one, - alloc_tensor_func=alloc_tensor_func, ) return output diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py index 35a649781d..5ccdbb4e9a 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/marlin_impl.py @@ -33,7 +33,6 @@ def _fused_experts( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func=torch.empty, ): if alpha is not None or limit is not None: raise NotImplementedError("FuseMoeMarlin does not support clamped SwiGLU") diff --git a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py index 83500d4eae..506b711fcf 100644 --- a/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py +++ b/lightllm/common/basemodel/layer_weights/meta_weights/fused_moe/impl/triton_impl.py @@ -67,7 +67,6 @@ def _fused_experts( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func=torch.empty, ): w13_weight, w13_scale = w13.weight, w13.weight_scale w2_weight, w2_scale = w2.weight, w2.weight_scale @@ -102,7 +101,6 @@ def fused_experts_with_topk( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func=torch.empty, ): return self._fused_experts( input_tensor=input_tensor, @@ -114,7 +112,6 @@ def fused_experts_with_topk( alpha=alpha, limit=limit, clamp_up_add_one=clamp_up_add_one, - alloc_tensor_func=alloc_tensor_func, ) def __call__( diff --git a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py index b26e8fa668..b9f5d23f1c 100644 --- a/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py +++ b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py @@ -78,18 +78,15 @@ def masked_group_gemm( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func: Callable = torch.empty, ): padded_m = recv_x[0].shape[1] E, N, _ = w1.shape block_size = 128 # groupgemm (masked layout) - gemm_out_a = alloc_tensor_func((E, padded_m, N), device=recv_x[0].device, dtype=dtype) + gemm_out_a = torch.empty((E, padded_m, N), device=recv_x[0].device, dtype=dtype) expected_m = min(expected_m, padded_m) - qsilu_out_scale = alloc_tensor_func( - (E, padded_m, N // 2 // block_size), device=recv_x[0].device, dtype=torch.float32 - ) - qsilu_out = alloc_tensor_func((E, padded_m, N // 2), dtype=w1.dtype, device=recv_x[0].device) + qsilu_out_scale = torch.empty((E, padded_m, N // 2 // block_size), device=recv_x[0].device, dtype=torch.float32) + qsilu_out = torch.empty((E, padded_m, N // 2), dtype=w1.dtype, device=recv_x[0].device) _deepgemm_grouped_fp8_nt_masked(recv_x, (w1, w1_scale), gemm_out_a, masked_m, expected_m) silu_and_mul_masked_post_quant_fwd( @@ -103,7 +100,7 @@ def masked_group_gemm( clamp_up_add_one=clamp_up_add_one, ) del gemm_out_a - gemm_out_b = alloc_tensor_func(recv_x[0].shape, device=recv_x[0].device, dtype=dtype) + gemm_out_b = torch.empty_like(recv_x[0], device=recv_x[0].device, dtype=dtype) _deepgemm_grouped_fp8_nt_masked((qsilu_out, qsilu_out_scale), (w2, w2_scale), gemm_out_b, masked_m, expected_m) return gemm_out_b @@ -219,7 +216,6 @@ def fused_experts( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func: Callable = torch.empty, ): assert (limit is None and alpha is None) or (limit is not None and alpha is not None) check_ep_expert_dtype(quant_method) @@ -247,7 +243,6 @@ def fused_experts( alpha=alpha, limit=limit, clamp_up_add_one=clamp_up_add_one, - alloc_tensor_func=alloc_tensor_func, ) @@ -269,7 +264,6 @@ def fused_experts_impl( alpha: Optional[float] = None, limit: Optional[float] = None, clamp_up_add_one: bool = True, - alloc_tensor_func: Callable = torch.empty, ): # Check constraints. assert hidden_states.shape[1] == w1.shape[2], "Hidden size mismatch" @@ -291,9 +285,7 @@ def fused_experts_impl( combined_x = None if is_prefill: - qinput_tensor, input_scale = per_token_group_quant_fp8( - hidden_states, block_size_k, dtype=w1.dtype, alloc_func=alloc_tensor_func - ) + qinput_tensor, input_scale = per_token_group_quant_fp8(hidden_states, block_size_k, dtype=w1.dtype) allocate_on_comm_stream = previous_event is not None # Expanded dispatch directly produces expert-contiguous, alignment-padded inputs: # recv_x[0]: [num_expanded_tokens, hidden] @@ -403,7 +395,6 @@ def fused_experts_impl( alpha=alpha, limit=limit, clamp_up_add_one=clamp_up_add_one, - alloc_tensor_func=alloc_tensor_func, ) # low latency combine combined_x, event_overlap, hook = buffer.low_latency_combine( diff --git a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py index a439a01f9d..32a1880f74 100644 --- a/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py +++ b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py @@ -156,8 +156,6 @@ def per_token_group_quant_fp8( dtype=torch.float32, ) - # Adapted from - # https://github.com/sgl-project/sglang/blob/7e257cd666c0d639626487987ea8e590da1e9395/python/sglang/srt/layers/quantization/fp8_kernel.py#L290 if HAS_SGL_KERNEL and not use_ue8m0_scales: finfo = torch.finfo(dtype) fp8_max, fp8_min = finfo.max, finfo.min diff --git a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py index 9e3c92d849..2f4e532428 100644 --- a/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py +++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py @@ -551,7 +551,6 @@ def _routed_experts( alpha=1.0, limit=float(self.swiglu_limit), clamp_up_add_one=False, - alloc_tensor_func=self.alloc_tensor, ) def _ffn_tp(self, input, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight): diff --git a/lightllm/server/core/objs/py_sampling_params.py b/lightllm/server/core/objs/py_sampling_params.py index 214bdf561d..ba3368d57c 100644 --- a/lightllm/server/core/objs/py_sampling_params.py +++ b/lightllm/server/core/objs/py_sampling_params.py @@ -4,7 +4,7 @@ """ import os from typing import List, Optional, Union, Tuple -from lightllm.utils.config_utils import get_generation_config_dict +from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF from .sampling_params import MAX_SEED @@ -111,11 +111,11 @@ def __init__( @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = get_generation_config_dict(weight_dir) + generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() def _cfg(key, default): - value = generation_cfg.get(key) - return value if value is not None else default + v = generation_cfg.get(key) + return v if v is not None else default cls._do_sample = _cfg("do_sample", False) cls._presence_penalty = _cfg("presence_penalty", 0.0) diff --git a/lightllm/server/core/objs/sampling_params.py b/lightllm/server/core/objs/sampling_params.py index 545d1876fb..f0dda1dd80 100644 --- a/lightllm/server/core/objs/sampling_params.py +++ b/lightllm/server/core/objs/sampling_params.py @@ -1,7 +1,7 @@ import os import ctypes from typing import Optional, List, Tuple, Union -from lightllm.utils.config_utils import get_generation_config_dict +from transformers import GenerationConfig from lightllm.server.req_id_generator import MAX_BEST_OF from lightllm.utils.envs_utils import get_env_start_args from .pd_kv_trans_params import PDKVTransParamObj @@ -406,11 +406,11 @@ def init(self, tokenizer, **kwargs): @classmethod def load_generation_cfg(cls, weight_dir): try: - generation_cfg = get_generation_config_dict(weight_dir) + generation_cfg = GenerationConfig.from_pretrained(weight_dir, trust_remote_code=True).to_dict() def _cfg(key, default): - value = generation_cfg.get(key) - return value if value is not None else default + v = generation_cfg.get(key) + return v if v is not None else default cls._do_sample = _cfg("do_sample", False) cls._presence_penalty = _cfg("presence_penalty", 0.0) diff --git a/revert.md b/revert.md index 343be7eb5b..c0de4d75e1 100644 --- a/revert.md +++ b/revert.md @@ -62,6 +62,7 @@ | Benchmark 与 autotune 产物 | 删除 static benchmark 增量及本 PR 新增的 H100/H200 autotune JSON | `388dd29de8a000500279ad273c06a6d7874f7cf1`、`4ac53e82fd3ecabef8ff0bbd4100a5fec3202821`、`0a48e27bcbfa8fec91f52bebb874616e1a2a95a0` | | DeepEP 环境自动调参 | 回退 decode dispatch capacity、NVSHMEM QP depth 等通用自动派生,恢复原有默认值 | `9a071ca5e0e53bab4bca63586f78e92b493cfb68`、`97ae2d12acdd77bcf87a786a710c0f1f049d8196`、`69c260f4a5f193fcd9170d06d0c77187a2ee18c1` | | B300 通用对齐 | 恢复设备判断;UE8M0 只在 SM100 上启用,不再无条件应用到所有设备 | `f5f3ed2cc74857ef8821840d5ca0d42c4f2a3e67` | +| 通用采样默认值优化 | 恢复 `core/objs` 直接通过 Transformers 加载完整 generation config,不再使用为改变默认 `top_k` 引入的共享 helper | `c19b8537134c66040a8dc1468c0b848650155967` (`default topk from huggingface's 50 to -1, if_inverse 70.6 -> 73`) | | DSV4 DP 结束 barrier review | 不纳入无条件 DSV4 barrier;保留此前 CPU-cache 场景已有的条件 barrier | `a2a7052d1a9ffdf765e81f7c43bf59807d484f08` | | mHC TileLang 启动预热 | 回退通用 `_kernel_warmup` hook、DSV4 mHC 预热流程及 split-K token 枚举;保留 MTP hidden 准备所需的 `hc_post` 引用 | `fc8b7ec411774fab269e1e3799efff4ac15f826e` (`warmup tilelang`) | From 6166eefbde1754c4140e8ac75f6fac783cb50821 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Tue, 29 Sep 2026 03:29:16 +0000 Subject: [PATCH 210/214] minor --- lightllm/common/quantization/deepgemm.py | 5 ++--- .../mode_backend/pd/decode_node_impl/decode_impl.py | 7 ++++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/lightllm/common/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py index 7ae8cad52f..6cdc3146d7 100644 --- a/lightllm/common/quantization/deepgemm.py +++ b/lightllm/common/quantization/deepgemm.py @@ -5,7 +5,6 @@ from lightllm.common.quantization.registry import QUANTMETHODS from lightllm.common.basemodel.triton_kernel.quantization.fp8act_quant_kernel import per_token_group_quant_fp8 from lightllm.utils.log_utils import init_logger -from lightllm.utils.device_utils import is_sm100_gpu logger = init_logger(__name__) @@ -63,7 +62,7 @@ def quantize(self, weight: torch.Tensor, output: WeightPack): from lightllm.common.basemodel.triton_kernel.quantization.fp8w8a8_block_quant_kernel import weight_quant device = output.weight.device - weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=is_sm100_gpu()) + weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=True) output.weight.copy_(weight) output.weight_scale.copy_(scale) return @@ -91,7 +90,7 @@ def apply( column_major_scales=True, scale_tma_aligned=True, alloc_func=alloc_func, - use_ue8m0_scales=is_sm100_gpu(), + use_ue8m0_scales=True, ) if out is None: diff --git a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py index e076cce8a0..237a423bda 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/decode_node_impl/decode_impl.py @@ -138,7 +138,10 @@ def _filter_not_ready_reqs(self, req_ids: List[int]) -> List[InferReq]: and not req_obj.infer_aborted and req_obj.cur_kv_len == req_obj.shm_req.input_len and req_obj.hybrid_cache_len > 0 - and req_obj.hybrid_cache_len == (req_obj.shm_req.input_len - 1) // 256 * 256 + and req_obj.hybrid_cache_len + == (req_obj.shm_req.input_len - 1) + // self.args.linear_att_hash_page_size + * self.args.linear_att_hash_page_size and req_obj.tail_small_page_buffer_id is None and self.radix_cache is not None ): @@ -201,8 +204,6 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq) -> bool: return False if self.radix_cache is not None: self.radix_cache.free_radix_cache_to_get_enough_token(need_mem_size) - if need_mem_size > mem_manager.allocator.can_use_mem_size: - return False prompt_page = req_manager.get_prompt_cache_page_size() swa_start = max(0, (input_len - 1) // prompt_page * prompt_page - prompt_page) From 6c818d2c12090484132090ee83c2f95eb6ae9d2e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 30 Sep 2026 01:17:23 +0000 Subject: [PATCH 211/214] minor --- .../mode_backend/pd/nixl_kv_transporter.py | 36 +++++++++---------- 1 file changed, 16 insertions(+), 20 deletions(-) diff --git a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py index cbfcc1ee4f..1cdd67eace 100644 --- a/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py +++ b/lightllm/server/router/model_infer/mode_backend/pd/nixl_kv_transporter.py @@ -84,24 +84,6 @@ def _create_paged_xfer_handles( descs = self.nixl_agent.get_xfer_descs(pages_data, "VRAM") return self.nixl_agent.prep_xfer_dlist(agent_name, descs, "VRAM") - def _get_local_page_xfer_handles(self, transfer_nbytes: int): - if transfer_nbytes not in self.page_local_xfer_handles: - self.page_local_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( - self.page_reg_desc, self.num_pages, transfer_nbytes - ) - return self.page_local_xfer_handles[transfer_nbytes] - - def _get_remote_page_xfer_handles(self, remote_agent: PDAgentMetadata, transfer_nbytes: int): - if transfer_nbytes not in remote_agent.page_xfer_handles: - page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) - remote_agent.page_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( - page_mem_desc, - remote_agent.num_pages, - transfer_nbytes, - agent_name=remote_agent.agent_name, - ) - return remote_agent.page_xfer_handles[transfer_nbytes] - def connect_add_remote_agent(self, remote_agent: PDAgentMetadata): with self._remote_agents_lock: if remote_agent.agent_name in self.remote_agents: @@ -292,8 +274,22 @@ def write_blocks_paged( assert trans_task.src_page_index is not None and trans_task.dst_page_index is not None assert trans_task.transfer_nbytes is not None remote_agent: PDAgentMetadata = self.remote_agents[decode_agent_name] - src_handle = self._get_local_page_xfer_handles(trans_task.transfer_nbytes) - dst_handle = self._get_remote_page_xfer_handles(remote_agent, trans_task.transfer_nbytes) + transfer_nbytes = trans_task.transfer_nbytes + if transfer_nbytes not in self.page_local_xfer_handles: + self.page_local_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( + self.page_reg_desc, self.num_pages, transfer_nbytes + ) + src_handle = self.page_local_xfer_handles[transfer_nbytes] + + if transfer_nbytes not in remote_agent.page_xfer_handles: + page_mem_desc = self.nixl_agent.deserialize_descs(remote_agent.page_reg_desc) + remote_agent.page_xfer_handles[transfer_nbytes] = self._create_paged_xfer_handles( + page_mem_desc, + remote_agent.num_pages, + transfer_nbytes, + agent_name=remote_agent.agent_name, + ) + dst_handle = remote_agent.page_xfer_handles[transfer_nbytes] handle = self.nixl_agent.make_prepped_xfer( "WRITE", src_handle, From 3b7a8b2e71d533a962ca63aa4489310a00023d4a Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 30 Sep 2026 01:43:15 +0000 Subject: [PATCH 212/214] delete dsv4 multi level kv cache --- .../deepseek4_mem_manager.py | 17 +- .../kv_cache_mem_manager/operator/deepseek.py | 97 ++++ lightllm/common/req_manager/deepseek4.py | 28 ++ .../model_infer/mode_backend/base_backend.py | 11 +- .../mode_backend/dsv4_multi_level_kv_cache.py | 472 ------------------ .../mode_backend/multi_level_kv_cache.py | 429 +++++++++++++++- revert.md | 8 + .../common/test_deepseek4_paged_cache.py | 142 +++++- .../deepseek_v4/test_vision_integration.py | 25 +- .../mode_backend/test_multi_level_kv_cache.py | 121 +++++ 10 files changed, 837 insertions(+), 513 deletions(-) delete mode 100644 lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py diff --git a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py index dcdbb2ecc3..5394ac78e7 100644 --- a/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py +++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py @@ -686,13 +686,19 @@ def get_loadable_cpu_cache_end( return loadable_end if loadable_end > loaded_start else 0 def prepare_cpu_cache_load( - self, token_num: int, loaded_end: int, resume_swa_slots: torch.Tensor + self, + token_num: int, + loaded_end: int, + resume_swa_slots: torch.Tensor, + mem_indexes: Optional[torch.Tensor] = None, ) -> DeepseekV4CpuCacheLoadPlan: - """Allocate a missing history suffix using the request's continuation slots. + """Prepare a missing history suffix using the request's continuation slots. ``loaded_end`` is the CPU checkpoint boundary. ``token_num`` may be smaller than the checkpoint page when a GPU radix prefix overlaps its beginning, but both endpoints remain 256-token aligned. + The common CPU-cache path supplies reserved slots; direct callers may + request their allocation here. """ token_num = int(token_num) loaded_end = int(loaded_end) @@ -710,9 +716,10 @@ def prepare_cpu_cache_load( block_num = token_num // DSV4_PROMPT_CACHE_PAGE_SIZE device = self.swa_pool.buffer.device - full_indexes_cpu = self.alloc(token_num) - - mem_indexes = full_indexes_cpu.to(device, non_blocking=True) + if mem_indexes is None: + mem_indexes = self.alloc(token_num).to(device, non_blocking=True) + else: + assert mem_indexes.numel() == token_num and mem_indexes.device == device history_full_slots = mem_indexes.view(block_num, DSV4_PROMPT_CACHE_PAGE_SIZE) history_c4_slots = history_full_slots[:, 3::4] // 4 if self.n_c4 else None history_c128_slots = history_full_slots[:, 127::128] // 128 if self.n_c128 else None diff --git a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py index 06b5640568..660b699f9d 100644 --- a/lightllm/common/kv_cache_mem_manager/operator/deepseek.py +++ b/lightllm/common/kv_cache_mem_manager/operator/deepseek.py @@ -1,3 +1,6 @@ +import dataclasses +from typing import List, Optional + import torch from .normal import NormalMemOperator from .base import BaseMemManagerOperator @@ -81,6 +84,10 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv: class DeepseekV4MemOperator(BaseMemManagerOperator): + def __init__(self, mem_manager): + super().__init__(mem_manager) + self.cpu_cache_staging_slots = [Dsv4StagingSlot(), Dsv4StagingSlot()] + def copy_mem_to_mem(self, src_mem_index: torch.Tensor, dst_mem_index: torch.Tensor): """Copy packed history pages; continuation is restored by the request manager.""" manager = self.mem_manager @@ -135,3 +142,93 @@ def load_cpu_cache_pages( first_page_history_offset_tokens, ) return + + def store_cpu_cache_pages( + self, + staging_slot: int, + source_mem_indexes: List[torch.Tensor], + source_req_meta: List[List[int]], + page_indexes: List[int], + cpu_cache_client, + producer_stream: torch.cuda.Stream, + cpu_stream: torch.cuda.Stream, + ): + """Pack into owned staging, then asynchronously scatter to CPU pages.""" + slot = self.cpu_cache_staging_slots[staging_slot] + assert not slot.in_use + page_num = len(page_indexes) + layout = self.mem_manager.cpu_cache_layout + if slot.buffer is None or slot.buffer.shape[0] < page_num: + with torch.cuda.stream(producer_stream): + slot.buffer = torch.empty((page_num, layout.page_nbytes), dtype=torch.uint8, device="cuda") + slot.source_mem_indexes = torch.empty( + (page_num, layout.token_page_size), dtype=torch.int32, device="cuda" + ) + slot.page_indexes_cuda = torch.empty((page_num,), dtype=torch.int32, device="cuda") + slot.page_indexes_cpu = torch.empty((page_num,), dtype=torch.int32, device="cpu", pin_memory=True) + slot.in_use = True + slot.page_indexes_cpu[:page_num].numpy()[:] = page_indexes + with torch.cuda.stream(producer_stream): + indexes = slot.source_mem_indexes[:page_num] + torch.stack(source_mem_indexes, out=indexes) + staging = slot.buffer[:page_num] + req_meta = torch.tensor(source_req_meta, dtype=torch.int32, device="cuda") + self.pack_cpu_cache_pages(indexes, req_meta, staging) + pack_event = torch.cuda.Event() + pack_event.record() + + # Request teardown and the next prefill may recycle the original slabs. + torch.cuda.current_stream().wait_event(pack_event) + with torch.cuda.stream(cpu_stream): + cpu_stream.wait_event(pack_event) + indexes_cuda = slot.page_indexes_cuda[:page_num] + indexes_cuda.copy_(slot.page_indexes_cpu[:page_num], non_blocking=True) + self.scatter_packed_cpu_cache_pages(staging, indexes_cuda, cpu_cache_client) + store_event = torch.cuda.Event() + store_event.record() + return pack_event, store_event + + def load_cpu_cache_to_gpu(self, mem_indexes, page_indexes, cpu_cache_client, req): + """Restore reserved history slots and request-owned continuation state.""" + from lightllm.server.router.model_infer.infer_batch import g_infer_context + + manager = self.mem_manager + req_manager = g_infer_context.req_manager + loaded_start = int(req.cur_kv_len) + loaded_end = loaded_start + mem_indexes.numel() + req_manager.prepare_swa(req.req_idx, loaded_end - 256, loaded_end) + resume_slots = req_manager.get_swa_slots( + req.req_idx, torch.arange(loaded_end - 256, loaded_end, device=mem_indexes.device) + ) + try: + plan = manager.prepare_cpu_cache_load( + token_num=mem_indexes.numel(), + loaded_end=loaded_end, + resume_swa_slots=resume_slots, + mem_indexes=mem_indexes, + ) + self.load_cpu_cache_pages( + plan=plan, + page_indexes=page_indexes, + cpu_cache_client=cpu_cache_client, + first_page_history_offset_tokens=loaded_start % manager.cpu_cache_layout.token_page_size, + ) + except Exception: + req_manager.clear_runtime_state(req.req_idx) + raise + + if g_infer_context.radix_cache is not None: + req_manager.restore_cpu_cache_checkpoints( + req, loaded_start, loaded_end, cpu_cache_client, g_infer_context.radix_cache + ) + + +@dataclasses.dataclass +class Dsv4StagingSlot: + """Operator-owned buffers leased until the cache module completes a store.""" + + buffer: Optional[torch.Tensor] = None + source_mem_indexes: Optional[torch.Tensor] = None + page_indexes_cpu: Optional[torch.Tensor] = None + page_indexes_cuda: Optional[torch.Tensor] = None + in_use: bool = False diff --git a/lightllm/common/req_manager/deepseek4.py b/lightllm/common/req_manager/deepseek4.py index 29d612e4df..43f462f5e5 100644 --- a/lightllm/common/req_manager/deepseek4.py +++ b/lightllm/common/req_manager/deepseek4.py @@ -137,6 +137,34 @@ def restore_state(self, req, state_cache_manager, buffer_idx, checkpoint_len): # A 256-token checkpoint closes both compressor groups. The next C128 # group overwrites all of its rows before reading them. + def restore_cpu_cache_checkpoints(self, req, loaded_start, loaded_end, cpu_cache_client, radix_cache): + """Retain loaded continuation checkpoints for subsequent GPU radix reuse.""" + from lightllm.utils.envs_utils import get_env_start_args + + args = get_env_start_args() + layout = self.mem_manager.cpu_cache_layout + page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() + big_page_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + first_boundary = (loaded_start // big_page_tokens + 1) * big_page_tokens + for boundary in range(first_boundary, loaded_end + 1, big_page_tokens): + buffer_idx = self.big_page_buffers.alloc_one_state_cache() + assert buffer_idx is not None + page_idx = page_list[boundary // layout.token_page_size - 1] + self.big_page_buffers.buffer[buffer_idx].copy_( + cpu_cache_client.cpu_kv_cache_tensor[page_idx, layout.swa_offset :] + ) + req.hybrid_len_to_big_page_id[boundary] = buffer_idx + if loaded_end == req.hybrid_cache_len and loaded_end % big_page_tokens: + radix_cache.free_one_small_page_buffer() + req.tail_small_page_buffer_id = self.small_page_buffers.alloc_one_state_cache() + if req.tail_small_page_buffer_id is not None: + self.save_state( + req.req_idx, + req.tail_small_page_buffer_id, + self.small_page_buffers, + checkpoint_len=loaded_end, + ) + def clear_runtime_state(self, req_idx): pages = self._swa_pages[req_idx] if pages: 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 0f17ab664e..47937e793a 100644 --- a/lightllm/server/router/model_infer/mode_backend/base_backend.py +++ b/lightllm/server/router/model_infer/mode_backend/base_backend.py @@ -50,7 +50,6 @@ create_cache_placement_controller, ) from .multi_level_kv_cache import MultiLevelKvCacheModule -from .dsv4_multi_level_kv_cache import Dsv4MultiLevelKvCacheModule from lightllm.utils.profiler import ProcessProfiler, ProfilerCmd @@ -258,8 +257,7 @@ def init_model(self, kvargs): self.init_spec_engine() if self.args.enable_cpu_cache: - cache_module_cls = Dsv4MultiLevelKvCacheModule if self.is_deepseek_v4 else MultiLevelKvCacheModule - self.multi_level_cache_module = cache_module_cls(self) + self.multi_level_cache_module = MultiLevelKvCacheModule(self) prof_name = f"lightllm-model_backend-node{self.node_rank}_dev{get_current_device_id()}" prof_mode = self.args.enable_profiling @@ -703,10 +701,7 @@ def _get_classed_reqs( # 定期对 radix cache 进行 merge,防止查询插入的操作效率下降 self._timer_merge_radix_tree() - if self.args.enable_cpu_cache and ( - (self.is_deepseek_v4 and self.is_master_in_dp) - or (not self.is_deepseek_v4 and len(g_infer_context.infer_req_ids) > 0) - ): + if self.args.enable_cpu_cache: self.multi_level_cache_module.update_cpu_cache_task_states() if req_ids is None: @@ -878,7 +873,7 @@ def _get_classed_reqs( # 一些可以复用的通用功能函数 def _pre_post_handle(self, run_reqs: List[InferReq], is_chuncked_mode: bool) -> List[InferReqUpdatePack]: update_func_objs: List[InferReqUpdatePack] = [] - cpu_store_reqs = [] if self.args.enable_cpu_cache and self.is_deepseek_v4 and self.is_master_in_dp else None + cpu_store_reqs = [] if self.args.enable_cpu_cache else None # 通用状态预先填充 is_master_in_dp = self.is_master_in_dp for req_obj in run_reqs: diff --git a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py deleted file mode 100644 index 34f998230a..0000000000 --- a/lightllm/server/router/model_infer/mode_backend/dsv4_multi_level_kv_cache.py +++ /dev/null @@ -1,472 +0,0 @@ -import bisect -import dataclasses -from collections import deque -from typing import Deque, Dict, List, Optional - -import torch -import torch.distributed as dist - -from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager -from lightllm.common.kv_cache_mem_manager.operator.deepseek import DeepseekV4MemOperator -from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuPageAllocState -from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context -from lightllm.utils.envs_utils import get_dsv4_cpu_cache_max_pages_per_task - -from .multi_level_kv_cache import MultiLevelKvCacheModule - - -def _split_dsv4_loaded_cache_lengths( - original_gpu_kv_len: int, - loaded_end: int, - requested_end: int, - disk_prompt_cache_len: int, -) -> tuple[int, int]: - """Split an actual CPU-cache load into CPU and disk matched token counts.""" - load_start = max(0, int(original_gpu_kv_len)) - load_end = max(load_start, int(loaded_end)) - actual_loaded_len = load_end - load_start - - matched_end = max(0, int(requested_end)) - matched_disk_len = min(max(0, int(disk_prompt_cache_len)), matched_end) - disk_start = matched_end - matched_disk_len - actual_disk_len = max(0, min(load_end, matched_end) - max(load_start, disk_start)) - actual_disk_len = min(actual_disk_len, actual_loaded_len) - actual_cpu_len = actual_loaded_len - actual_disk_len - return actual_cpu_len, actual_disk_len - - -class Dsv4MultiLevelKvCacheModule(MultiLevelKvCacheModule): - def __init__(self, backend): - super().__init__(backend) - self._dsv4_store_sessions: Dict[int, Dsv4CpuStoreSession] = {} - self._dsv4_store_tasks: Deque[Dsv4StoreTask] = deque() - self._dsv4_staging_slots = [Dsv4StagingSlot(), Dsv4StagingSlot()] - self._dsv4_max_pages_per_store_task = get_dsv4_cpu_cache_max_pages_per_task() - - def _try_release_dsv4_session(self, session: "Dsv4CpuStoreSession") -> None: - if not session.closing or not session.load_submitted or session.pending_task_num != 0: - return - if session.load_event is not None and not session.load_event.query(): - return - - if session.leased_pages: - self.cpu_cache_client.lock.acquire_sleep1ms() - try: - if self.args.enable_disk_cache: - # A disk-cache group must contain one complete request prefix in - # root-to-tail order. Incremental store batches may complete in - # a different order, so do not publish or release any page until - # every page leased by this session is ready. - if not self.cpu_cache_client.check_allpages_ready(session.leased_pages): - return - self.cpu_cache_client.update_pages_status_to_ready( - page_list=session.leased_pages, - deref=True, - disk_offload_enable=True, - token_num_in_page_list=(len(session.leased_pages) * self.args.cpu_cache_token_page_size), - ) - else: - # Cumulative hashes make the root page the most valuable entry. - # Releasing tail-to-root makes the tail oldest in the LRU. - self.cpu_cache_client.deref_pages(list(reversed(session.leased_pages))) - finally: - self.cpu_cache_client.lock.release() - del self._dsv4_store_sessions[session.request_id] - - def _poll_dsv4_store_tasks(self, wait_for_one: bool = False) -> None: - if not self._dsv4_store_tasks: - for session in list(self._dsv4_store_sessions.values()): - self._try_release_dsv4_session(session) - return - - completed = [] - if wait_for_one: - self._dsv4_store_tasks[0].store_event.synchronize() - while self._dsv4_store_tasks and self._dsv4_store_tasks[0].store_event.query(): - completed.append(self._dsv4_store_tasks.popleft()) - if completed: - self.cpu_cache_client.lock.acquire_sleep1ms() - try: - for task in completed: - self.cpu_cache_client.update_pages_status_to_ready(task.owner_pages, deref=False) - finally: - self.cpu_cache_client.lock.release() - - touched_sessions = {} - for task in completed: - slot = self._dsv4_staging_slots[task.staging_slot] - slot.in_use = False - for session in task.sessions: - session.pending_task_num -= 1 - assert session.pending_task_num >= 0 - touched_sessions[session.request_id] = session - for session in touched_sessions.values(): - self._try_release_dsv4_session(session) - for session in list(self._dsv4_store_sessions.values()): - self._try_release_dsv4_session(session) - - def _submit_dsv4_store_batch( - self, - store_pages: List["Dsv4StorePage"], - producer_stream: torch.cuda.Stream, - ) -> None: - assert 0 < len(store_pages) <= self._dsv4_max_pages_per_store_task - sessions = {item.session.request_id: item.session for item in store_pages} - owner_pages = [item.cpu_page_index for item in store_pages] - operator: DeepseekV4MemOperator = self.backend.model.mem_manager.operator - cpu_stream = g_infer_context.get_cpu_kv_cache_stream() - - self._poll_dsv4_store_tasks() - slot_index = None - while slot_index is None: - for candidate, slot in enumerate(self._dsv4_staging_slots): - if not slot.in_use: - slot_index = candidate - break - if slot_index is None: - self._poll_dsv4_store_tasks(wait_for_one=True) - - slot = self._dsv4_staging_slots[slot_index] - page_num = len(store_pages) - page_nbytes = self.backend.model.mem_manager.cpu_cache_layout.page_nbytes - if slot.buffer is None or slot.buffer.shape[0] < page_num: - with torch.cuda.stream(producer_stream): - buffer = torch.empty((page_num, page_nbytes), dtype=torch.uint8, device="cuda") - source_mem_indexes = torch.empty( - (page_num, self.backend.model.mem_manager.cpu_cache_layout.token_page_size), - dtype=torch.int32, - device="cuda", - ) - page_indexes_cuda = torch.empty((page_num,), dtype=torch.int32, device="cuda") - page_indexes_cpu = torch.empty((page_num,), dtype=torch.int32, device="cpu", pin_memory=True) - slot.buffer = buffer - slot.source_mem_indexes = source_mem_indexes - slot.page_indexes_cuda = page_indexes_cuda - slot.page_indexes_cpu = page_indexes_cpu - slot.in_use = True - - slot.page_indexes_cpu[:page_num].numpy()[:] = owner_pages - with torch.cuda.stream(producer_stream): - source_mem_indexes = slot.source_mem_indexes[:page_num] - torch.stack([item.source_mem_indexes for item in store_pages], out=source_mem_indexes) - staging = slot.buffer[:page_num] - source_req_meta = torch.tensor( - [[item.req_idx, item.checkpoint_len] for item in store_pages], dtype=torch.int32, device="cuda" - ) - operator.pack_cpu_cache_pages(source_mem_indexes, source_req_meta, staging) - pack_event = torch.cuda.Event() - pack_event.record() - - # Request free/pause runs on the current stream. It may recycle the - # original DS4 slabs after packing, but never before it. - torch.cuda.current_stream().wait_event(pack_event) - with torch.cuda.stream(cpu_stream): - cpu_stream.wait_event(pack_event) - page_indexes_cuda = slot.page_indexes_cuda[:page_num] - page_indexes_cuda.copy_(slot.page_indexes_cpu[:page_num], non_blocking=True) - operator.scatter_packed_cpu_cache_pages(staging, page_indexes_cuda, self.cpu_cache_client) - store_event = torch.cuda.Event() - store_event.record() - - for session in sessions.values(): - session.pending_task_num += 1 - self._dsv4_store_tasks.append( - Dsv4StoreTask( - owner_pages=owner_pages, - sessions=list(sessions.values()), - staging_slot=slot_index, - pack_event=pack_event, - store_event=store_event, - ) - ) - - def store_completed_prefill_pages( - self, - reqs: List[InferReq], - producer_stream: torch.cuda.Stream, - ) -> None: - """Incrementally snapshot newly completed DS4 checkpoints before source reuse.""" - layout = self.backend.model.mem_manager.cpu_cache_layout - token_page_size = layout.token_page_size - store_pages: List[Dsv4StorePage] = [] - closing_sessions = {} - self.cpu_cache_client.lock.acquire_sleep1ms() - try: - for req in reqs: - session = self._dsv4_store_sessions.get(req.req_id) - if session is None or session.closing: - continue - token_hashes = req.shm_req.token_hash_list.get_all() - if session.disabled: - closing_sessions[session.request_id] = session - continue - if session.next_page_index >= len(token_hashes): - closing_sessions[session.request_id] = session - continue - page_lens = req.shm_req.token_hash_page_len_list.get_all() - target_page_index = bisect.bisect_right(page_lens, req.cur_kv_len) - if target_page_index <= session.next_page_index: - continue - - start_page_index = session.next_page_index - page_indexes, alloc_states = self.cpu_cache_client.allocate_pages( - token_hashes[start_page_index:target_page_index], - disk_offload_enable=False, - ) - for offset, (cpu_page_index, alloc_state) in enumerate(zip(page_indexes, alloc_states)): - if cpu_page_index == -1: - session.disabled = True - break - checkpoint_index = start_page_index + offset - session.leased_pages.append(cpu_page_index) - session.next_page_index += 1 - if alloc_state is CpuPageAllocState.NEW_STORE_OWNER: - token_start = checkpoint_index * token_page_size - source_mem_indexes = self.backend.model.req_manager.req_to_token_indexs[ - req.req_idx, token_start : token_start + token_page_size - ] - store_pages.append( - Dsv4StorePage( - session=session, - cpu_page_index=cpu_page_index, - source_mem_indexes=source_mem_indexes, - req_idx=req.req_idx, - checkpoint_len=token_start + token_page_size, - ) - ) - if session.disabled or session.next_page_index >= len(token_hashes): - closing_sessions[session.request_id] = session - finally: - self.cpu_cache_client.lock.release() - - for offset in range(0, len(store_pages), self._dsv4_max_pages_per_store_task): - self._submit_dsv4_store_batch( - store_pages[offset : offset + self._dsv4_max_pages_per_store_task], - producer_stream=producer_stream, - ) - for session in closing_sessions.values(): - session.closing = True - self._try_release_dsv4_session(session) - - @staticmethod - def _get_image_safe_load_end(req: InferReq, loaded_start: int, load_end: int, page_size: int) -> int: - """Move an image-internal CPU resume point before that image.""" - for image_start, image_end in reversed(req.image_block_spans): - if image_start < load_end < image_end: - load_end = image_start // page_size * page_size - return load_end if load_end > loaded_start else 0 - - def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): - idle_token_num = g_infer_context.get_can_alloc_token_num() - is_master_in_dp = self.backend.is_master_in_dp - for req in reqs: - page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() - page_len_list = req.shm_req.token_hash_page_len_list.get_all() - assert len(page_list) <= len(page_len_list) - - gpu_kv_len = int(req.cur_kv_len) - requested_end = gpu_kv_len - matched_disk_len = int(req.shm_req.disk_prompt_cache_len) - if is_master_in_dp: - session = Dsv4CpuStoreSession( - request_id=req.req_id, - next_page_index=len(page_list), - leased_pages=list(page_list), - ) - page_size = self.backend.model.mem_manager.cpu_cache_layout.token_page_size - # A radix-owned checkpoint without its CPU prefix creates an unreachable hash-chain hole. - if gpu_kv_len // page_size > session.next_page_index: - session.disabled = True - session.closing = True - self._dsv4_store_sessions[req.req_id] = session - - loaded_end = gpu_kv_len - if page_list: - mem_manager: DeepseekV4MemoryManager = self.backend.model.mem_manager - layout = mem_manager.cpu_cache_layout - requested_end = int(page_len_list[len(page_list) - 1]) - if requested_end > gpu_kv_len: - swa_capacity = g_infer_context.get_can_alloc_dsv4_swa_page_num() - loadable_end = mem_manager.get_loadable_cpu_cache_end( - gpu_kv_len, - requested_end, - idle_token_num, - swa_capacity, - ) - loadable_end = self._get_image_safe_load_end(req, gpu_kv_len, loadable_end, layout.token_page_size) - if loadable_end != 0: - token_num = loadable_end - gpu_kv_len - full_need = token_num - if self.backend.radix_cache is not None: - radix_cache = self.backend.radix_cache - radix_cache.free_radix_cache_to_get_enough_token(full_need) - - loadable_end = mem_manager.get_loadable_cpu_cache_end( - gpu_kv_len, - loadable_end, - int(mem_manager.allocator.can_use_mem_size), - int(mem_manager.swa_page_allocator.can_use_mem_size), - ) - loadable_end = self._get_image_safe_load_end( - req, gpu_kv_len, loadable_end, layout.token_page_size - ) - if loadable_end != 0: - loaded_end = loadable_end - token_num = loaded_end - gpu_kv_len - first_page_index = gpu_kv_len // layout.token_page_size - cpu_pages = page_list[first_page_index : loaded_end // layout.token_page_size] - page_indexes_cuda = torch.tensor(cpu_pages, dtype=torch.int32, device="cuda") - req_manager = self.backend.model.req_manager - req_manager.prepare_swa(req.req_idx, loaded_end - 256, loaded_end) - resume_slots = req_manager.get_swa_slots( - req.req_idx, torch.arange(loaded_end - 256, loaded_end, device="cuda") - ) - plan = mem_manager.prepare_cpu_cache_load( - token_num=token_num, loaded_end=loaded_end, resume_swa_slots=resume_slots - ) - try: - mem_manager.operator.load_cpu_cache_pages( - plan=plan, - page_indexes=page_indexes_cuda, - cpu_cache_client=self.cpu_cache_client, - first_page_history_offset_tokens=gpu_kv_len % layout.token_page_size, - ) - except Exception: - mem_manager.free(plan.mem_indexes) - req_manager.clear_runtime_state(req.req_idx) - raise - self.backend.model.req_manager.req_to_token_indexs[ - req.req_idx, gpu_kv_len:loaded_end - ] = plan.mem_indexes - req.cur_kv_len = loaded_end - req.hold_kv_len = loaded_end - if self.backend.radix_cache is not None: - big_page_tokens = ( - self.args.linear_att_hash_page_size * self.args.linear_att_page_block_num - ) - first_boundary = (gpu_kv_len // big_page_tokens + 1) * big_page_tokens - for boundary in range(first_boundary, loaded_end + 1, big_page_tokens): - buffer_idx = mem_manager.big_page_buffers.alloc_one_state_cache() - assert buffer_idx is not None - page_idx = page_list[boundary // layout.token_page_size - 1] - mem_manager.big_page_buffers.buffer[buffer_idx].copy_( - self.cpu_cache_client.cpu_kv_cache_tensor[page_idx, layout.swa_offset :] - ) - req.hybrid_len_to_big_page_id[boundary] = buffer_idx - if loaded_end == req.hybrid_cache_len and loaded_end % big_page_tokens: - self.backend.radix_cache.free_one_small_page_buffer() - req.tail_small_page_buffer_id = ( - req_manager.small_page_buffers.alloc_one_state_cache() - ) - if req.tail_small_page_buffer_id is not None: - req_manager.save_state( - req.req_idx, - req.tail_small_page_buffer_id, - req_manager.small_page_buffers, - checkpoint_len=loaded_end, - ) - idle_token_num -= token_num - - if is_master_in_dp: - cpu_prompt_cache_len, disk_prompt_cache_len = _split_dsv4_loaded_cache_lengths( - original_gpu_kv_len=gpu_kv_len, - loaded_end=loaded_end, - requested_end=requested_end, - disk_prompt_cache_len=matched_disk_len, - ) - req.shm_req.cpu_prompt_cache_len = cpu_prompt_cache_len - req.shm_req.disk_prompt_cache_len = disk_prompt_cache_len - req.shm_req.shm_cur_kv_len = loaded_end - session.load_submitted = True - if loaded_end > gpu_kv_len: - session.load_event = torch.cuda.Event() - session.load_event.record() - - dist.barrier(group=self.init_sync_group) - if is_master_in_dp: - for session in list(self._dsv4_store_sessions.values()): - self._try_release_dsv4_session(session) - return - - def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> List[InferReq]: - if self.backend.is_master_in_dp: - for req in finished_reqs: - session = self._dsv4_store_sessions.get(req.req_id) - if session is not None: - session.closing = True - self._try_release_dsv4_session(session) - self._poll_dsv4_store_tasks() - # Source pages are fenced by the pack event. Request teardown does - # not wait for the independent staging-to-host transfer. - return finished_reqs - - def update_cpu_cache_task_states(self): - self._poll_dsv4_store_tasks() - return - - -@dataclasses.dataclass -class Dsv4CpuStoreSession: - """跟踪单个 DS4 请求持有的 CPU pages 及其异步 load/store 生命周期。""" - - request_id: int - # [0, next_page_index) 的 checkpoint pages 已处理;该值指向下一个待处理 page。 - next_page_index: int = 0 - # 本 session 持有引用的 CPU page 编号,session 释放时统一 deref。 - leased_pages: List[int] = dataclasses.field(default_factory=list) - # 已提交但尚未完成的 GPU -> CPU store batch 数量。 - pending_task_num: int = 0 - # 当前 hash 链无法继续存储,不再预留新的 CPU page。 - disabled: bool = False - # 不再接收新的 store page,等待 load/store 完成后释放 session。 - closing: bool = False - # 本请求的初始 CPU -> GPU load 流程已经提交。 - load_submitted: bool = False - # 初始 CPU -> GPU load 的完成事件;没有实际 load 时为 None。 - load_event: Optional[torch.cuda.Event] = None - - -@dataclasses.dataclass -class Dsv4StorePage: - """描述一个由当前请求负责写入的GPU page -> CPU page""" - - # 持有该 CPU page 引用并跟踪异步任务的请求 session。 - session: Dsv4CpuStoreSession - # 已预留、等待写入的目标 CPU page 编号。 - cpu_page_index: int - # 该 checkpoint page 对应的 GPU KV slot 编号。 - source_mem_indexes: torch.Tensor - req_idx: int - checkpoint_len: int - - -@dataclasses.dataclass -class Dsv4StoreTask: - """跟踪一个已经提交的异步 GPU -> CPU store batch。""" - - # 本 batch 负责写入的 CPU pages,完成后统一发布为 READY。 - owner_pages: List[int] - # 本 batch 涉及的请求 session,完成后分别减少 pending_task_num。 - sessions: List[Dsv4CpuStoreSession] - # 本 batch 占用的 staging slot 编号。 - staging_slot: int - # GPU KV 已打包完成;此事件完成后原始 KV slot 可以被回收。 - pack_event: torch.cuda.Event - # staging 数据已写入 CPU;轮询此事件判断整个 batch 是否完成。 - store_event: torch.cuda.Event - - -@dataclasses.dataclass -class Dsv4StagingSlot: - """可复用的 GPU staging buffer 及其 CPU page 索引缓冲区。""" - - # 打包后的连续 GPU 字节缓冲区,形状为 [page_capacity, page_nbytes]。 - buffer: Optional[torch.Tensor] = None - # batch 内每个 checkpoint page 对应的 GPU KV slot 编号。 - source_mem_indexes: Optional[torch.Tensor] = None - # Python 写入的 pinned CPU page 编号,用于异步拷贝到 GPU。 - page_indexes_cpu: Optional[torch.Tensor] = None - # scatter kernel 使用的目标 CPU page 编号。 - page_indexes_cuda: Optional[torch.Tensor] = None - # True 表示该 slot 仍被一个未完成的 store task 占用。 - in_use: bool = False diff --git a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py index d6a02fa19e..dd5cd970fc 100644 --- a/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py +++ b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py @@ -4,17 +4,18 @@ import dataclasses import bisect from functools import lru_cache -from typing import Optional, List, Deque +from typing import Optional, List, Deque, Dict from collections import deque from lightllm.server.multi_level_kv_cache import CacheTier from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuKvCacheClient, CpuPageAllocState from lightllm.utils.config_utils import is_hybrid_att_model -from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.envs_utils import get_env_start_args, get_dsv4_cpu_cache_max_pages_per_task from ..infer_batch import InferReq from lightllm.utils.dist_utils import create_new_group_for_current_dp from lightllm.common.basemodel.triton_kernel.kv_cache_offload import offload_gpu_kv_to_cpu, load_cpu_kv_to_gpu from lightllm.server.router.model_infer.infer_batch import g_infer_context from lightllm.utils.log_utils import init_logger +from lightllm.common.kv_cache_mem_manager.operator.deepseek import DeepseekV4MemOperator logger = init_logger(__name__) @@ -39,6 +40,10 @@ def __init__(self, backend): self.cpu_cache_handle_queue: Deque[TransTask] = deque() self.cpu_cache_client = CpuKvCacheClient(only_create_meta_data=False, init_shm_data=False) + if isinstance(self.backend.model.mem_manager.operator, DeepseekV4MemOperator): + self._dsv4_store_sessions: Dict[int, Dsv4CpuStoreSession] = {} + self._dsv4_store_tasks: Deque[Dsv4StoreTask] = deque() + self._dsv4_max_pages_per_store_task = get_dsv4_cpu_cache_max_pages_per_task() @lru_cache() def need_sync_compute_stream(self) -> bool: @@ -61,20 +66,42 @@ def need_sync_compute_stream(self) -> bool: return False def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): - idle_token_num = g_infer_context.get_can_alloc_token_num() - all_page_list = [] + cache_reqs = [] + pages_to_release = [] is_master_in_dp = self.backend.is_master_in_dp + is_deepseek_v4 = g_infer_context.is_deepseek_v4 for req in reqs: - page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() - # 需要返回 prompt logprobs 的请求不应加载 cpu cache: - # 命中后会复用缓存 kv、跳过推理,拿不到对应 logprobs。 - # match 侧通常已跳过;这里仍要 deref 已 match 的 page,避免引用泄漏。 - if req.sampling_param.shm_param.prompt_logprobs >= 0: + # KV 命中会跳过计算,无法返回对应的 prompt logprobs。 + # match 侧通常已跳过;这里仍需释放之前匹配的页面引用。 + skip_cpu_cache = req.sampling_param.shm_param.prompt_logprobs >= 0 + if skip_cpu_cache: if is_master_in_dp: req.shm_req.cpu_prompt_cache_len = 0 - all_page_list.extend(page_list) - continue + req.shm_req.disk_prompt_cache_len = 0 + else: + cache_reqs.append(req) + # DSV4 正常加载的页面由 session 持有,等待异步 load/store 完成。 + if is_master_in_dp and (skip_cpu_cache or not is_deepseek_v4): + pages_to_release.extend(req.shm_req.cpu_cache_match_page_indexes.get_all()) + + if is_deepseek_v4: + self._load_dsv4_cpu_cache_to_reqs(cache_reqs) + else: + self._load_standard_cpu_cache_to_reqs(cache_reqs) + + if is_master_in_dp and pages_to_release: + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + self.cpu_cache_client.deref_pages(pages_to_release) + finally: + self.cpu_cache_client.lock.release() + return + def _load_standard_cpu_cache_to_reqs(self, reqs: List[InferReq]): + idle_token_num = g_infer_context.get_can_alloc_token_num() + is_master_in_dp = self.backend.is_master_in_dp + for req in reqs: + page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() page_len_list = req.shm_req.token_hash_page_len_list.get_all() page_len_start_list = [0] + page_len_list assert len(page_list) <= len(page_len_list) @@ -148,14 +175,7 @@ def load_cpu_cache_to_reqs(self, reqs: List[InferReq]): if self.backend.is_master_in_dp: req.shm_req.shm_cur_kv_len = req.cur_kv_len - all_page_list.extend(page_list) - dist.barrier(group=self.init_sync_group) - - if self.backend.is_master_in_dp: - self.cpu_cache_client.lock.acquire_sleep1ms() - self.cpu_cache_client.deref_pages(page_list=all_page_list) - self.cpu_cache_client.lock.release() return def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> List[InferReq]: @@ -164,6 +184,8 @@ def offload_finished_reqs_to_cpu_cache(self, finished_reqs: List[InferReq]) -> L """ # 如果开启了cpu cache,将达到finished状态的请求开启将gpu kv cache 卸载到 cpu cache中的操作。 # 当 kv cache 卸载完成后,才会进行请求的真实退出操作。 + if g_infer_context.is_deepseek_v4: + return self._finish_dsv4_cpu_cache_sessions(finished_reqs) true_finished_reqs = [] cpu_stream = g_infer_context.get_cpu_kv_cache_stream() for req in finished_reqs: @@ -327,6 +349,12 @@ def _handle_hybrid_att_last_page(self, req: InferReq, move_block_size: int, page return move_block_size def update_cpu_cache_task_states(self): + if g_infer_context.is_deepseek_v4: + if self.backend.is_master_in_dp: + self._poll_dsv4_store_tasks() + return + if not g_infer_context.infer_req_ids: + return if self.backend.is_master_in_dp: trans_ok_tasks = [] while len(self.cpu_cache_handle_queue) != 0: @@ -362,6 +390,300 @@ def update_cpu_cache_task_states(self): task.req_obj.cpu_cache_task_status = InferReq._CpuCacheTaskStatus.FINISHED return + def _try_release_dsv4_session(self, session: "Dsv4CpuStoreSession") -> None: + if not session.closing or not session.load_submitted or session.pending_task_num != 0: + return + if session.load_event is not None and not session.load_event.query(): + return + + if session.leased_pages: + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + if self.args.enable_disk_cache: + # A disk-cache group must contain one complete request prefix in + # root-to-tail order. Incremental store batches may complete in + # a different order, so do not publish or release any page until + # every page leased by this session is ready. + if not self.cpu_cache_client.check_allpages_ready(session.leased_pages): + return + self.cpu_cache_client.update_pages_status_to_ready( + page_list=session.leased_pages, + deref=True, + disk_offload_enable=True, + token_num_in_page_list=(len(session.leased_pages) * self.args.cpu_cache_token_page_size), + ) + else: + # Cumulative hashes make the root page the most valuable entry. + # Releasing tail-to-root makes the tail oldest in the LRU. + self.cpu_cache_client.deref_pages(list(reversed(session.leased_pages))) + finally: + self.cpu_cache_client.lock.release() + del self._dsv4_store_sessions[session.request_id] + + def _poll_dsv4_store_tasks(self, wait_for_one: bool = False) -> None: + if not self._dsv4_store_tasks: + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + return + + completed = [] + if wait_for_one: + self._dsv4_store_tasks[0].store_event.synchronize() + while self._dsv4_store_tasks and self._dsv4_store_tasks[0].store_event.query(): + completed.append(self._dsv4_store_tasks.popleft()) + if completed: + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + for task in completed: + self.cpu_cache_client.update_pages_status_to_ready(task.owner_pages, deref=False) + finally: + self.cpu_cache_client.lock.release() + + touched_sessions = {} + for task in completed: + slot = self.backend.model.mem_manager.operator.cpu_cache_staging_slots[task.staging_slot] + slot.in_use = False + for session in task.sessions: + session.pending_task_num -= 1 + assert session.pending_task_num >= 0 + touched_sessions[session.request_id] = session + for session in touched_sessions.values(): + self._try_release_dsv4_session(session) + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + + def _submit_dsv4_store_batch( + self, + store_pages: List["Dsv4StorePage"], + producer_stream: torch.cuda.Stream, + ) -> None: + assert 0 < len(store_pages) <= self._dsv4_max_pages_per_store_task + sessions = {item.session.request_id: item.session for item in store_pages} + owner_pages = [item.cpu_page_index for item in store_pages] + operator: DeepseekV4MemOperator = self.backend.model.mem_manager.operator + cpu_stream = g_infer_context.get_cpu_kv_cache_stream() + + self._poll_dsv4_store_tasks() + slot_index = None + while slot_index is None: + for candidate, slot in enumerate(operator.cpu_cache_staging_slots): + if not slot.in_use: + slot_index = candidate + break + if slot_index is None: + self._poll_dsv4_store_tasks(wait_for_one=True) + + pack_event, store_event = operator.store_cpu_cache_pages( + staging_slot=slot_index, + source_mem_indexes=[item.source_mem_indexes for item in store_pages], + source_req_meta=[[item.req_idx, item.checkpoint_len] for item in store_pages], + page_indexes=owner_pages, + cpu_cache_client=self.cpu_cache_client, + producer_stream=producer_stream, + cpu_stream=cpu_stream, + ) + + for session in sessions.values(): + session.pending_task_num += 1 + self._dsv4_store_tasks.append( + Dsv4StoreTask( + owner_pages=owner_pages, + sessions=list(sessions.values()), + staging_slot=slot_index, + pack_event=pack_event, + store_event=store_event, + ) + ) + + def store_completed_prefill_pages( + self, + reqs: List[InferReq], + producer_stream: torch.cuda.Stream, + ) -> None: + """Incrementally snapshot newly completed DS4 checkpoints before source reuse.""" + if not g_infer_context.is_deepseek_v4 or not self.backend.is_master_in_dp: + return + layout = self.backend.model.mem_manager.cpu_cache_layout + token_page_size = layout.token_page_size + store_pages: List[Dsv4StorePage] = [] + closing_sessions = {} + self.cpu_cache_client.lock.acquire_sleep1ms() + try: + for req in reqs: + session = self._dsv4_store_sessions.get(req.req_id) + if session is None or session.closing: + continue + token_hashes = req.shm_req.token_hash_list.get_all() + if session.disabled: + closing_sessions[session.request_id] = session + continue + if session.next_page_index >= len(token_hashes): + closing_sessions[session.request_id] = session + continue + page_lens = req.shm_req.token_hash_page_len_list.get_all() + target_page_index = bisect.bisect_right(page_lens, req.cur_kv_len) + if target_page_index <= session.next_page_index: + continue + + start_page_index = session.next_page_index + page_indexes, alloc_states = self.cpu_cache_client.allocate_pages( + token_hashes[start_page_index:target_page_index], + disk_offload_enable=False, + ) + for offset, (cpu_page_index, alloc_state) in enumerate(zip(page_indexes, alloc_states)): + if cpu_page_index == -1: + session.disabled = True + break + checkpoint_index = start_page_index + offset + session.leased_pages.append(cpu_page_index) + session.next_page_index += 1 + if alloc_state is CpuPageAllocState.NEW_STORE_OWNER: + token_start = checkpoint_index * token_page_size + source_mem_indexes = self.backend.model.req_manager.req_to_token_indexs[ + req.req_idx, token_start : token_start + token_page_size + ] + store_pages.append( + Dsv4StorePage( + session=session, + cpu_page_index=cpu_page_index, + source_mem_indexes=source_mem_indexes, + req_idx=req.req_idx, + checkpoint_len=token_start + token_page_size, + ) + ) + if session.disabled or session.next_page_index >= len(token_hashes): + closing_sessions[session.request_id] = session + finally: + self.cpu_cache_client.lock.release() + + for offset in range(0, len(store_pages), self._dsv4_max_pages_per_store_task): + self._submit_dsv4_store_batch( + store_pages[offset : offset + self._dsv4_max_pages_per_store_task], + producer_stream=producer_stream, + ) + for session in closing_sessions.values(): + session.closing = True + self._try_release_dsv4_session(session) + + @staticmethod + def _get_image_safe_load_end(req: InferReq, loaded_start: int, load_end: int, page_size: int) -> int: + """Move an image-internal CPU resume point before that image.""" + for image_start, image_end in reversed(req.image_block_spans): + if image_start < load_end < image_end: + load_end = image_start // page_size * page_size + return load_end if load_end > loaded_start else 0 + + def _load_dsv4_cpu_cache_to_reqs(self, reqs: List[InferReq]): + idle_token_num = g_infer_context.get_can_alloc_token_num() + is_master_in_dp = self.backend.is_master_in_dp + for req in reqs: + page_list = req.shm_req.cpu_cache_match_page_indexes.get_all() + page_len_list = req.shm_req.token_hash_page_len_list.get_all() + assert len(page_list) <= len(page_len_list) + + gpu_kv_len = int(req.cur_kv_len) + requested_end = gpu_kv_len + matched_disk_len = int(req.shm_req.disk_prompt_cache_len) + if is_master_in_dp: + session = Dsv4CpuStoreSession( + request_id=req.req_id, + next_page_index=len(page_list), + leased_pages=list(page_list), + ) + page_size = self.backend.model.mem_manager.cpu_cache_layout.token_page_size + # A radix-owned checkpoint without its CPU prefix creates an unreachable hash-chain hole. + if gpu_kv_len // page_size > session.next_page_index: + session.disabled = True + session.closing = True + self._dsv4_store_sessions[req.req_id] = session + + loaded_end = gpu_kv_len + if page_list: + mem_manager = self.backend.model.mem_manager + layout = mem_manager.cpu_cache_layout + requested_end = int(page_len_list[len(page_list) - 1]) + if requested_end > gpu_kv_len: + swa_capacity = g_infer_context.get_can_alloc_dsv4_swa_page_num() + loadable_end = mem_manager.get_loadable_cpu_cache_end( + gpu_kv_len, + requested_end, + idle_token_num, + swa_capacity, + ) + loadable_end = self._get_image_safe_load_end(req, gpu_kv_len, loadable_end, layout.token_page_size) + if loadable_end != 0: + token_num = loadable_end - gpu_kv_len + full_need = token_num + if self.backend.radix_cache is not None: + radix_cache = self.backend.radix_cache + radix_cache.free_radix_cache_to_get_enough_token(full_need) + + loadable_end = mem_manager.get_loadable_cpu_cache_end( + gpu_kv_len, + loadable_end, + int(mem_manager.allocator.can_use_mem_size), + int(mem_manager.swa_page_allocator.can_use_mem_size), + ) + loadable_end = self._get_image_safe_load_end( + req, gpu_kv_len, loadable_end, layout.token_page_size + ) + if loadable_end != 0: + loaded_end = loadable_end + token_num = loaded_end - gpu_kv_len + first_page_index = gpu_kv_len // layout.token_page_size + cpu_pages = page_list[first_page_index : loaded_end // layout.token_page_size] + page_indexes_cuda = torch.tensor(cpu_pages, dtype=torch.int32, device="cuda") + mem_indexes = mem_manager.alloc(token_num).cuda(non_blocking=True) + try: + mem_manager.operator.load_cpu_cache_to_gpu( + mem_indexes=mem_indexes, + page_indexes=page_indexes_cuda, + cpu_cache_client=self.cpu_cache_client, + req=req, + ) + except Exception: + mem_manager.free(mem_indexes) + raise + self.backend.model.req_manager.req_to_token_indexs[ + req.req_idx, gpu_kv_len:loaded_end + ] = mem_indexes + req.cur_kv_len = loaded_end + req.hold_kv_len = loaded_end + idle_token_num -= token_num + + if is_master_in_dp: + cpu_prompt_cache_len, disk_prompt_cache_len = _split_dsv4_loaded_cache_lengths( + original_gpu_kv_len=gpu_kv_len, + loaded_end=loaded_end, + requested_end=requested_end, + disk_prompt_cache_len=matched_disk_len, + ) + req.shm_req.cpu_prompt_cache_len = cpu_prompt_cache_len + req.shm_req.disk_prompt_cache_len = disk_prompt_cache_len + req.shm_req.shm_cur_kv_len = loaded_end + session.load_submitted = True + if loaded_end > gpu_kv_len: + session.load_event = torch.cuda.Event() + session.load_event.record() + + dist.barrier(group=self.init_sync_group) + if is_master_in_dp: + for session in list(self._dsv4_store_sessions.values()): + self._try_release_dsv4_session(session) + return + + def _finish_dsv4_cpu_cache_sessions(self, finished_reqs: List[InferReq]) -> List[InferReq]: + if self.backend.is_master_in_dp: + for req in finished_reqs: + session = self._dsv4_store_sessions.get(req.req_id) + if session is not None: + session.closing = True + self._try_release_dsv4_session(session) + self._poll_dsv4_store_tasks() + # Source pages are fenced by the pack event. Request teardown does + # not wait for the independent staging-to-host transfer. + return finished_reqs + @dataclasses.dataclass class TransTask: @@ -370,3 +692,74 @@ class TransTask: page_readies: torch.Tensor req_obj: InferReq sync_event: torch.cuda.Event + + +def _split_dsv4_loaded_cache_lengths( + original_gpu_kv_len: int, + loaded_end: int, + requested_end: int, + disk_prompt_cache_len: int, +) -> tuple[int, int]: + """Split an actual CPU-cache load into CPU and disk matched token counts.""" + load_start = max(0, int(original_gpu_kv_len)) + load_end = max(load_start, int(loaded_end)) + actual_loaded_len = load_end - load_start + + matched_end = max(0, int(requested_end)) + matched_disk_len = min(max(0, int(disk_prompt_cache_len)), matched_end) + disk_start = matched_end - matched_disk_len + actual_disk_len = max(0, min(load_end, matched_end) - max(load_start, disk_start)) + actual_disk_len = min(actual_disk_len, actual_loaded_len) + actual_cpu_len = actual_loaded_len - actual_disk_len + return actual_cpu_len, actual_disk_len + + +@dataclasses.dataclass +class Dsv4CpuStoreSession: + """跟踪单个 DS4 请求持有的 CPU pages 及其异步 load/store 生命周期。""" + + request_id: int + # [0, next_page_index) 的 checkpoint pages 已处理;该值指向下一个待处理 page。 + next_page_index: int = 0 + # 本 session 持有引用的 CPU page 编号,session 释放时统一 deref。 + leased_pages: List[int] = dataclasses.field(default_factory=list) + # 已提交但尚未完成的 GPU -> CPU store batch 数量。 + pending_task_num: int = 0 + # 当前 hash 链无法继续存储,不再预留新的 CPU page。 + disabled: bool = False + # 不再接收新的 store page,等待 load/store 完成后释放 session。 + closing: bool = False + # 本请求的初始 CPU -> GPU load 流程已经提交。 + load_submitted: bool = False + # 初始 CPU -> GPU load 的完成事件;没有实际 load 时为 None。 + load_event: Optional[torch.cuda.Event] = None + + +@dataclasses.dataclass +class Dsv4StorePage: + """描述一个由当前请求负责写入的GPU page -> CPU page""" + + # 持有该 CPU page 引用并跟踪异步任务的请求 session。 + session: Dsv4CpuStoreSession + # 已预留、等待写入的目标 CPU page 编号。 + cpu_page_index: int + # 该 checkpoint page 对应的 GPU KV slot 编号。 + source_mem_indexes: torch.Tensor + req_idx: int + checkpoint_len: int + + +@dataclasses.dataclass +class Dsv4StoreTask: + """跟踪一个已经提交的异步 GPU -> CPU store batch。""" + + # 本 batch 负责写入的 CPU pages,完成后统一发布为 READY。 + owner_pages: List[int] + # 本 batch 涉及的请求 session,完成后分别减少 pending_task_num。 + sessions: List[Dsv4CpuStoreSession] + # 本 batch 占用的 staging slot 编号。 + staging_slot: int + # GPU KV 已打包完成;此事件完成后原始 KV slot 可以被回收。 + pack_event: torch.cuda.Event + # staging 数据已写入 CPU;轮询此事件判断整个 batch 是否完成。 + store_event: torch.cuda.Event diff --git a/revert.md b/revert.md index c0de4d75e1..01a16ef7b5 100644 --- a/revert.md +++ b/revert.md @@ -67,3 +67,11 @@ | mHC TileLang 启动预热 | 回退通用 `_kernel_warmup` hook、DSV4 mHC 预热流程及 split-K token 枚举;保留 MTP hidden 准备所需的 `hc_post` 引用 | `fc8b7ec411774fab269e1e3799efff4ac15f826e` (`warmup tilelang`) | 清理后,`lightllm/server/httpserver_for_pd_master/`、`lightllm/server/router/req_queue/`、`lightllm/server/multi_level_kv_cache/disk_cache_worker.py` 和 `lightllm/utils/device_utils.py` 相对拆分基线无额外差异;`communication_op.py` 只保留 DSV4 `experts_` 字段兼容,HTTP server 只保留 Vision 图像块不可跨 prefill 切分的校验。此前确认的 MTP CUDA Graph hidden 输入修复继续保留。 + +## 2026-09-30 多级缓存架构收敛 + +删除独立的 `Dsv4MultiLevelKvCacheModule`,统一通过 main 的 `MultiLevelKvCacheModule` 管理 CPU 页面引用、异步任务及磁盘发布。保留 DSV4 必需的 prefill 增量 checkpoint 保存;staging 与 pack/unpack 回归 `DeepseekV4MemOperator`,SWA 和 radix checkpoint 恢复回归 `DeepseekV4ReqManager`。保留原有 CPU 页布局及两阶段完成事件,不改变 PD 传输协议。 + +prompt-logprobs 请求的过滤、命中长度清零和匹配页引用释放统一放在公共加载入口;普通模型和 DSV4 仅在实际缓存加载阶段分派,不再分别实现过滤策略。 + +普通模型的正常匹配页也在公共入口统一释放,加载子流程保留结束 barrier;DSV4 的正常匹配页继续由异步 session 持有,公共入口只释放其跳过加载的页面。 diff --git a/unit_tests/common/test_deepseek4_paged_cache.py b/unit_tests/common/test_deepseek4_paged_cache.py index 9823eff337..2161d3f87a 100644 --- a/unit_tests/common/test_deepseek4_paged_cache.py +++ b/unit_tests/common/test_deepseek4_paged_cache.py @@ -234,6 +234,140 @@ def test_cpu_cache_roundtrip_uses_derived_history_slots(cache): assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages +@pytest.mark.parametrize("gpu_prefix_len", [0, 256]) +def test_common_cpu_cache_incrementally_preserves_recycled_checkpoints(cache, monkeypatch, gpu_prefix_len): + from collections import deque + from lightllm.server.multi_level_kv_cache.cpu_cache_client import CpuPageAllocState + from lightllm.server.router.model_infer.mode_backend import multi_level_kv_cache as module + + manager, requests = cache + req_idx, src = _reserve(manager, requests, 4096) + cpu_pages = torch.empty((2, manager.cpu_cache_layout.page_nbytes), dtype=torch.uint8, pin_memory=True) + published = [] + released = [] + client = SimpleNamespace( + cpu_kv_cache_tensor=cpu_pages, + lock=SimpleNamespace(acquire_sleep1ms=lambda: None, release=lambda: None), + allocate_pages=lambda hashes, **kwargs: (list(hashes), [CpuPageAllocState.NEW_STORE_OWNER] * len(hashes)), + update_pages_status_to_ready=lambda pages, **kwargs: published.extend(pages), + deref_pages=released.extend, + ) + cache_module = module.MultiLevelKvCacheModule.__new__(module.MultiLevelKvCacheModule) + cache_module.backend = SimpleNamespace( + is_master_in_dp=True, radix_cache=None, model=SimpleNamespace(mem_manager=manager, req_manager=requests) + ) + cache_module.args = SimpleNamespace(enable_disk_cache=False) + cache_module.cpu_cache_client = client + cache_module.init_sync_group = None + cache_module._dsv4_store_sessions = {} + cache_module._dsv4_store_tasks = deque() + cache_module._dsv4_max_pages_per_store_task = 1 + cpu_stream = torch.cuda.Stream() + monkeypatch.setattr(module.g_infer_context, "is_deepseek_v4", True) + monkeypatch.setattr(module.g_infer_context, "req_manager", requests) + monkeypatch.setattr(module.g_infer_context, "radix_cache", None) + monkeypatch.setattr(module.g_infer_context, "get_cpu_kv_cache_stream", lambda: cpu_stream) + monkeypatch.setattr(module.g_infer_context, "get_can_alloc_token_num", lambda: manager.allocator.can_use_mem_size) + monkeypatch.setattr( + module.g_infer_context, "get_can_alloc_dsv4_swa_page_num", lambda: manager.swa_page_allocator.can_use_mem_size + ) + monkeypatch.setattr(module.dist, "barrier", lambda group: None) + req = SimpleNamespace( + req_id=1, + req_idx=req_idx, + cur_kv_len=0, + hold_kv_len=4096, + image_block_spans=[], + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), + shm_req=SimpleNamespace( + cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: []), + token_hash_list=SimpleNamespace(get_all=lambda: [0, 1]), + token_hash_page_len_list=SimpleNamespace(get_all=lambda: [2048, 4096]), + disk_prompt_cache_len=0, + ), + ) + cache_module.load_cpu_cache_to_reqs([req]) + expected_history = [] + for start, end in ((0, 2048), (2048, 4096)): + requests.prepare_swa(req_idx, start, end) + for pool in (manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool, manager.swa_pool): + pool.buffer.fill_(11 + start // 2048) + manager.c4_state_buffer.fill_(21 + start // 2048) + manager.c4_indexer_state_buffer.fill_(31 + start // 2048) + expected_history.append( + [ + pool.read(0, src[start + ratio - 1 : end : ratio].long() // ratio).clone() + for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)) + ] + ) + req.cur_kv_len = end + with torch.cuda.stream(cpu_stream): + torch.cuda._sleep(20_000_000) + cache_module.store_completed_prefill_pages([req], torch.cuda.current_stream()) + + # Free the source while the independent staging-to-CPU writes may still run. + assert cache_module.offload_finished_reqs_to_cpu_cache([req]) == [req] + requests.free([req_idx], src) + torch.cuda.synchronize() + cache_module.update_cpu_cache_task_states() + assert sorted(published) == [0, 1] + assert released == [1, 0] + assert cache_module._dsv4_store_sessions == {} + assert not any(slot.in_use for slot in manager.operator.cpu_cache_staging_slots) + + req.req_idx = requests.alloc() + req.req_id = 2 + req.cur_kv_len = req.hold_kv_len = gpu_prefix_len + req.hybrid_cache_len = 4096 + req.hybrid_len_to_big_page_id = {} + if gpu_prefix_len: + prefix = manager.alloc(gpu_prefix_len).cuda() + requests.req_to_token_indexs[req.req_idx, :gpu_prefix_len] = prefix + for index, (pool, ratio) in enumerate( + ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)) + ): + pool.write( + 0, prefix[ratio - 1 :: ratio].long() // ratio, expected_history[0][index][: gpu_prefix_len // ratio] + ) + radix = SimpleNamespace(free_radix_cache_to_get_enough_token=lambda tokens: None) + cache_module.backend.radix_cache = radix + monkeypatch.setattr(module.g_infer_context, "radix_cache", radix) + req.shm_req.cpu_cache_match_page_indexes.get_all = lambda: [0, 1] + cache_module.load_cpu_cache_to_reqs([req]) + assert req.cur_kv_len == req.hold_kv_len == 4096 + assert req.shm_req.cpu_prompt_cache_len == 4096 - gpu_prefix_len + assert list(req.hybrid_len_to_big_page_id) == [2048, 4096] + for boundary, index in req.hybrid_len_to_big_page_id.items(): + torch.testing.assert_close( + manager.big_page_buffers.buffer[index], + cpu_pages[boundary // 2048 - 1, manager.cpu_cache_layout.swa_offset :], + ) + restored = requests.req_to_token_indexs[req.req_idx, :4096].clone() + for page, (start, end) in enumerate(((0, 2048), (2048, 4096))): + for index, (pool, ratio) in enumerate( + ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)) + ): + torch.testing.assert_close( + pool.read(0, restored[start + ratio - 1 : end : ratio].long() // ratio), + expected_history[page][index], + rtol=0, + atol=0, + ) + swa_slots = requests.get_swa_slots(req.req_idx, torch.arange(3840, 4096, device="cuda")).long() + for layer in range(manager.layer_num): + assert manager.swa_pool.read(layer, swa_slots).eq(12).all() + tail_rows = swa_slots[-4:] // 128 * manager.c4_state_ring + swa_slots[-4:] % manager.c4_state_ring + assert manager.c4_state_buffer[:, tail_rows].eq(22).all() + assert manager.c4_indexer_state_buffer[:, tail_rows].eq(32).all() + cache_module.offload_finished_reqs_to_cpu_cache([req]) + torch.cuda.synchronize() + cache_module.update_cpu_cache_task_states() + manager.big_page_buffers.free_state_cache(list(req.hybrid_len_to_big_page_id.values())) + requests.free([req.req_idx], restored) + assert manager.allocator.can_use_mem_size == manager.size + assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages + + @pytest.mark.parametrize("length", [256, 512, 2048]) def test_shared_history_restores_private_continuation(cache, length): manager, requests = cache @@ -578,11 +712,11 @@ def request(req_idx, total_len): def test_cpu_load_failure_releases_reserved_history_and_private_swa(cache, monkeypatch): - from lightllm.server.router.model_infer.mode_backend import dsv4_multi_level_kv_cache as module + from lightllm.server.router.model_infer.mode_backend import multi_level_kv_cache as module manager, requests = cache req_idx = requests.alloc() - cache_module = module.Dsv4MultiLevelKvCacheModule.__new__(module.Dsv4MultiLevelKvCacheModule) + cache_module = module.MultiLevelKvCacheModule.__new__(module.MultiLevelKvCacheModule) cache_module.backend = SimpleNamespace( is_master_in_dp=False, radix_cache=None, model=SimpleNamespace(mem_manager=manager, req_manager=requests) ) @@ -594,6 +728,7 @@ def test_cpu_load_failure_releases_reserved_history_and_private_swa(cache, monke cur_kv_len=0, hold_kv_len=0, image_block_spans=[], + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), shm_req=SimpleNamespace( cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [0]), token_hash_page_len_list=SimpleNamespace(get_all=lambda: [2048]), @@ -605,6 +740,9 @@ def fail(**kwargs): raise RuntimeError("injected copy failure") monkeypatch.setattr(manager.operator, "load_cpu_cache_pages", fail) + monkeypatch.setattr(module.g_infer_context, "is_deepseek_v4", True) + monkeypatch.setattr(module.g_infer_context, "req_manager", requests) + monkeypatch.setattr(module.g_infer_context, "radix_cache", None) monkeypatch.setattr(module.g_infer_context, "get_can_alloc_token_num", lambda: manager.allocator.can_use_mem_size) monkeypatch.setattr( module.g_infer_context, "get_can_alloc_dsv4_swa_page_num", lambda: manager.swa_page_allocator.can_use_mem_size diff --git a/unit_tests/models/deepseek_v4/test_vision_integration.py b/unit_tests/models/deepseek_v4/test_vision_integration.py index e4ef71ca18..1c1e0f98af 100644 --- a/unit_tests/models/deepseek_v4/test_vision_integration.py +++ b/unit_tests/models/deepseek_v4/test_vision_integration.py @@ -245,18 +245,18 @@ def test_recover_swa_budget_includes_atomic_image_block(monkeypatch): ], ) def test_cpu_cache_load_end_never_splits_an_image(loaded_start, load_end, spans, expected): - from lightllm.server.router.model_infer.mode_backend.dsv4_multi_level_kv_cache import ( - Dsv4MultiLevelKvCacheModule, + from lightllm.server.router.model_infer.mode_backend.multi_level_kv_cache import ( + MultiLevelKvCacheModule, ) req = SimpleNamespace(image_block_spans=spans) - assert Dsv4MultiLevelKvCacheModule._get_image_safe_load_end(req, loaded_start, load_end, 2048) == expected + assert MultiLevelKvCacheModule._get_image_safe_load_end(req, loaded_start, load_end, 2048) == expected def test_cpu_cache_rechecks_image_boundary_after_capacity_changes(monkeypatch): from lightllm.server.router.model_infer.mode_backend import ( - dsv4_multi_level_kv_cache as cache_module, + multi_level_kv_cache as cache_module, ) capacity_results = iter([6144, 4096]) @@ -269,9 +269,9 @@ def get_loadable_cpu_cache_end(*args): capacity_calls.append(args) return next(capacity_results) - def prepare_cpu_cache_load(*, token_num, loaded_end, resume_swa_slots): + def prepare_cpu_cache_load(*, token_num, loaded_end, resume_swa_slots, mem_indexes): prepare_calls.append((token_num, loaded_end)) - return SimpleNamespace(mem_indexes=torch.arange(token_num, dtype=torch.int32)) + return SimpleNamespace(mem_indexes=mem_indexes) def load_cpu_cache_pages(*, page_indexes, **kwargs): loaded_pages.append(page_indexes.tolist()) @@ -290,9 +290,13 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): swa_page_allocator=SimpleNamespace(can_use_mem_size=2), get_loadable_cpu_cache_end=get_loadable_cpu_cache_end, prepare_cpu_cache_load=prepare_cpu_cache_load, - operator=SimpleNamespace(load_cpu_cache_pages=load_cpu_cache_pages), + alloc=lambda token_num: torch.arange(token_num, dtype=torch.int32), ) - module = object.__new__(cache_module.Dsv4MultiLevelKvCacheModule) + from lightllm.common.kv_cache_mem_manager.operator.deepseek import DeepseekV4MemOperator + + mem_manager.operator = DeepseekV4MemOperator(mem_manager) + mem_manager.operator.load_cpu_cache_pages = load_cpu_cache_pages + module = object.__new__(cache_module.MultiLevelKvCacheModule) module.backend = SimpleNamespace( is_master_in_dp=True, radix_cache=None, @@ -307,6 +311,7 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): req_idx=0, cur_kv_len=0, image_block_spans=[(3500, 4500)], + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), shm_req=SimpleNamespace( cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [10, 11, 12]), token_hash_page_len_list=SimpleNamespace(get_all=lambda: [2048, 4096, 6144]), @@ -323,7 +328,11 @@ def load_cpu_cache_pages(*, page_indexes, **kwargs): lambda data, **kwargs: real_tensor(data, **{key: value for key, value in kwargs.items() if key != "device"}), ) monkeypatch.setattr(cache_module.torch.cuda, "Event", lambda: SimpleNamespace(record=lambda: None)) + monkeypatch.setattr(torch.Tensor, "cuda", lambda self, **kwargs: self) monkeypatch.setattr(cache_module.dist, "barrier", lambda group: None) + monkeypatch.setattr(cache_module.g_infer_context, "is_deepseek_v4", True) + monkeypatch.setattr(cache_module.g_infer_context, "req_manager", req_manager) + monkeypatch.setattr(cache_module.g_infer_context, "radix_cache", None) monkeypatch.setattr(cache_module.g_infer_context, "get_can_alloc_token_num", lambda: 8192) monkeypatch.setattr( cache_module.g_infer_context, diff --git a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py index 260dbe1065..ad01bcf507 100644 --- a/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py +++ b/unit_tests/server/router/model_infer/mode_backend/test_multi_level_kv_cache.py @@ -200,6 +200,7 @@ def alloc_req_kv_mem(req, alloc_token_num): module.need_sync_compute_stream = lambda: False module.cpu_cache_client = SimpleNamespace() context = SimpleNamespace( + is_deepseek_v4=False, req_manager=SimpleNamespace(req_to_token_indexs=table, mem_manager=SimpleNamespace(alloc=alloc)), get_can_alloc_token_num=lambda: 200, ) @@ -226,3 +227,123 @@ def alloc_req_kv_mem(req, alloc_token_num): assert req.hold_kv_len == 132 assert loaded_indexes == list(range(132)) assert table[0, :132].tolist() == list(range(132)) + + +@pytest.mark.parametrize("disk_cache", [False, True]) +def test_dsv4_store_survives_request_finish_and_waits_for_load(disk_cache, monkeypatch): + module = MultiLevelKvCacheModule.__new__(MultiLevelKvCacheModule) + slot = SimpleNamespace(in_use=True) + module.backend = SimpleNamespace( + is_master_in_dp=True, + model=SimpleNamespace(mem_manager=SimpleNamespace(operator=SimpleNamespace(cpu_cache_staging_slots=[slot]))), + ) + module.args = SimpleNamespace(enable_disk_cache=disk_cache, cpu_cache_token_page_size=2048) + load_event = SimpleNamespace(ready=False) + load_event.query = lambda: load_event.ready + store_event = SimpleNamespace(ready=False) + store_event.query = lambda: store_event.ready + session = multi_level_kv_cache_impl.Dsv4CpuStoreSession( + request_id=3, leased_pages=[5, 7], pending_task_num=1, load_submitted=True, load_event=load_event + ) + module._dsv4_store_sessions = {3: session} + module._dsv4_store_tasks = deque( + [multi_level_kv_cache_impl.Dsv4StoreTask([7], [session], 0, object(), store_event)] + ) + published = [] + released = [] + module.cpu_cache_client = SimpleNamespace( + lock=SimpleNamespace(acquire_sleep1ms=lambda: None, release=lambda: None), + update_pages_status_to_ready=lambda page_list, **kwargs: published.append((list(page_list), kwargs)), + deref_pages=lambda pages: released.extend(pages), + check_allpages_ready=lambda pages: True, + ) + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "is_deepseek_v4", True) + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "infer_req_ids", []) + req = SimpleNamespace(req_id=3) + + assert module.offload_finished_reqs_to_cpu_cache([req]) == [req] + assert session.closing and session.pending_task_num == 1 + assert published == released == [] + assert slot.in_use + + store_event.ready = True + module.update_cpu_cache_task_states() + assert published == [([7], {"deref": False})] + assert released == [] and 3 in module._dsv4_store_sessions + assert session.pending_task_num == 0 and not slot.in_use + + load_event.ready = True + module.update_cpu_cache_task_states() + assert module._dsv4_store_sessions == {} + if disk_cache: + assert published[-1] == ([5, 7], {"deref": True, "disk_offload_enable": True, "token_num_in_page_list": 4096}) + else: + assert released == [7, 5] + + +@pytest.mark.parametrize("is_deepseek_v4", [False, True]) +@pytest.mark.parametrize("is_master_in_dp", [False, True]) +@pytest.mark.parametrize("mixed_batch", [False, True]) +def test_prompt_logprobs_filter_is_shared_by_models(is_deepseek_v4, is_master_in_dp, mixed_batch, monkeypatch): + module = MultiLevelKvCacheModule.__new__(MultiLevelKvCacheModule) + module.backend = SimpleNamespace(is_master_in_dp=is_master_in_dp) + loaded = [] + released = [] + module._load_dsv4_cpu_cache_to_reqs = loaded.extend + module._load_standard_cpu_cache_to_reqs = loaded.extend + module.cpu_cache_client = SimpleNamespace( + lock=SimpleNamespace(acquire_sleep1ms=lambda: None, release=lambda: None), + deref_pages=released.extend, + ) + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "is_deepseek_v4", is_deepseek_v4) + req = SimpleNamespace( + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=0)), + shm_req=SimpleNamespace( + cpu_prompt_cache_len=2048, + disk_prompt_cache_len=2048, + cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [4, 5]), + ), + ) + cache_req = SimpleNamespace( + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=-1)), + shm_req=SimpleNamespace(cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda: [7])), + ) + module.load_cpu_cache_to_reqs([req, cache_req] if mixed_batch else [req]) + assert loaded == ([cache_req] if mixed_batch else []) + expected_pages = [4, 5, 7] if mixed_batch and not is_deepseek_v4 else [4, 5] + assert released == (expected_pages if is_master_in_dp else []) + assert req.shm_req.cpu_prompt_cache_len == req.shm_req.disk_prompt_cache_len == (0 if is_master_in_dp else 2048) + + +def test_standard_load_releases_skipped_and_matched_pages_once_after_barrier(monkeypatch): + module = MultiLevelKvCacheModule.__new__(MultiLevelKvCacheModule) + module.backend = SimpleNamespace(is_master_in_dp=True) + module.init_sync_group = object() + events = [] + module.cpu_cache_client = SimpleNamespace( + lock=SimpleNamespace( + acquire_sleep1ms=lambda: events.append("lock"), + release=lambda: events.append("unlock"), + ), + deref_pages=lambda pages: events.append(("deref", list(pages))), + ) + reqs = [ + SimpleNamespace( + cur_kv_len=0, + sampling_param=SimpleNamespace(shm_param=SimpleNamespace(prompt_logprobs=logprobs)), + shm_req=SimpleNamespace( + input_len=64, + disk_prompt_cache_len=0, + cpu_cache_match_page_indexes=SimpleNamespace(get_all=lambda page=page: [page]), + token_hash_page_len_list=SimpleNamespace(get_all=lambda: [64]), + ), + ) + for page, logprobs in ((4, 0), (7, -1)) + ] + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "is_deepseek_v4", False) + monkeypatch.setattr(multi_level_kv_cache_impl.g_infer_context, "get_can_alloc_token_num", lambda: 128) + monkeypatch.setattr(multi_level_kv_cache_impl.dist, "barrier", lambda group: events.append("barrier")) + + module.load_cpu_cache_to_reqs(reqs) + + assert events == ["barrier", "lock", ("deref", [4, 7]), "unlock"] From 15c8c201ad058d0ff73e26e3cd33a22a2c409f2e Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 30 Sep 2026 01:54:23 +0000 Subject: [PATCH 213/214] refactor(mtp): unify async CPU accept lengths and ModelInput mirror updates --- lightllm/common/basemodel/batch_objs.py | 11 +++ .../mode_backend/chunked_prefill/impl.py | 5 +- .../mode_backend/dp_backend/impl.py | 15 ++-- .../mtp_speculative/dp_overlap_engine.py | 15 ++-- .../dp_overlap_proposers/eagle_with_att.py | 34 ++----- .../model_infer/mtp_speculative/engine.py | 11 +-- .../mtp_speculative/proposers/dspark.py | 2 - .../proposers/eagle_with_att.py | 22 ++--- revert.md | 4 + .../mtp_speculative/test_dspark.py | 9 +- .../mtp_speculative/test_eagle_overlap.py | 84 ++++++++++++++++- .../mtp_speculative/test_eagle_with_att.py | 89 ++++++++++++++++++- 12 files changed, 224 insertions(+), 77 deletions(-) diff --git a/lightllm/common/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py index 5f2f3a587c..93ccf9aeb4 100644 --- a/lightllm/common/basemodel/batch_objs.py +++ b/lightllm/common/basemodel/batch_objs.py @@ -74,6 +74,17 @@ def capture_cpu_mirrors(self): self._capture_cpu_mirror("b_ready_cache_len", "b_ready_cache_len_cpu") return + def select_mtp_cpu_mirrors(self, accept_len: torch.Tensor) -> None: + """Select accepted tails on a draft copy and advance to its first decode token.""" + req_start_rows = torch.nonzero(self.b_mtp_index_cpu == 0, as_tuple=False).flatten() + accepted_tail_rows = req_start_rows + accept_len - 1 + self.b_req_idx_cpu = self.b_req_idx_cpu.index_select(0, accepted_tail_rows) + self.b_mtp_index_cpu = torch.zeros_like(self.b_req_idx_cpu, dtype=self.b_mtp_index_cpu.dtype) + self.b_seq_len_cpu = self.b_seq_len_cpu.index_select(0, accepted_tail_rows) + 1 + + def advance_cpu_seq_len(self) -> None: + self.b_seq_len_cpu.add_(1) + def to_cuda(self): self.check_input() 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 97515faf0d..5abd9ce0f8 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 @@ -11,7 +11,7 @@ ) from lightllm.server.router.model_infer.mode_backend.generic_post_process import sample from lightllm.server.router.model_infer.infer_batch import g_infer_context -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine from lightllm.server.router.model_infer.mtp_speculative import utils as mtp_utils from lightllm.utils.log_utils import init_logger @@ -310,8 +310,7 @@ def decode_mtp( b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=spec_plan.draft_step, accept_len=mtp_accept_len, - accept_len_cpu=mtp_accept_len_cpu, - accept_len_ready_event=verify_event, + accept_len_cpu=AsyncPinnedCpuTensor(tensor=mtp_accept_len_cpu, ready_event=verify_event), ) mtp_utils.scatter_mtp_next_tokens( backend=self, 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 13f9feee7d..663a24e6c9 100644 --- a/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py +++ b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py @@ -15,7 +15,7 @@ from lightllm.server.router.model_infer.mode_backend.overlap_events import OverlapEventPack from lightllm.utils.dist_utils import get_current_device_id from lightllm.utils.envs_utils import get_env_start_args -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager 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 @@ -559,8 +559,9 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]): b_req_mtp_start_loc=b_req_mtp_start_loc, draft_step=spec_plan.draft_step, accept_len=mtp_accept_len, - accept_len_cpu=mtp_accept_len_cpu, - accept_len_ready_event=verify_event, + accept_len_cpu=( + AsyncPinnedCpuTensor(tensor=mtp_accept_len_cpu, ready_event=verify_event) if req_num > 0 else None + ), ) if req_num > 0: mtp_utils.scatter_mtp_next_tokens( @@ -765,8 +766,6 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf logits0 = model_output0.logits logits1 = model_output1.logits run_reqs = run_reqs0 + run_reqs1 - mtp_accept_len_cpu0 = None - mtp_accept_len_cpu1 = None if req_num > 0: assert len(run_reqs) == verify_row_num logits = torch.empty( @@ -831,13 +830,13 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf target_model_output0=model_output0, target_next_token_ids0=target_next_token_ids0, accept_len0=mtp_accept_len0, - accept_len_cpu0=mtp_accept_len_cpu0, target_model_input1=model_input1, target_model_output1=model_output1, target_next_token_ids1=target_next_token_ids1, accept_len1=mtp_accept_len1, - accept_len_cpu1=mtp_accept_len_cpu1, - accept_len_ready_event=verify_event, + accept_len_cpu=( + AsyncPinnedCpuTensor(tensor=mtp_accept_len_cpu, ready_event=verify_event) if req_num > 0 else None + ), draft_step=spec_plan.draft_step, ) if req_num > 0: diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py index e364fc928c..08d1a1a2a8 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_engine.py @@ -11,6 +11,9 @@ from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.base import ( BaseDpOverlapProposer, ) +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_proposers.eagle_with_att import ( + DpOverlapEagleWithAttProposer, +) 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 ( @@ -150,9 +153,7 @@ def propose_next_overlap( target_next_token_ids1: torch.Tensor, # [verify_batch_size1] accept_len1: torch.Tensor, # [real_req_num1] draft_step: int, - accept_len_cpu0: Optional[torch.Tensor] = None, # [real_req_num0] - accept_len_cpu1: Optional[torch.Tensor] = None, # [real_req_num1] - accept_len_ready_event: Optional[torch.cuda.Event] = None, + accept_len_cpu: Optional[AsyncPinnedCpuTensor] = None, # [real_req_num0 + real_req_num1] ) -> SpecProposal: assert target_next_token_ids0.shape == (target_model_input0.batch_size,) assert target_next_token_ids1.shape == (target_model_input1.batch_size,) @@ -160,12 +161,8 @@ def propose_next_overlap( assert accept_len1.ndim == 1 proposer_kwargs = {} - if self.proposer.backend.is_deepseek_v4: - proposer_kwargs = { - "accept_len_cpu0": accept_len_cpu0, - "accept_len_cpu1": accept_len_cpu1, - "accept_len_ready_event": accept_len_ready_event, - } + if self.proposer.backend.is_deepseek_v4 and isinstance(self.proposer, DpOverlapEagleWithAttProposer): + proposer_kwargs = {"accept_len_cpu": accept_len_cpu} return self.proposer.propose_next_overlap( target_model_input0=target_model_input0, target_model_output0=target_model_output0, diff --git a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py index db774f62aa..e0dc86a34d 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/dp_overlap_proposers/eagle_with_att.py @@ -11,7 +11,7 @@ get_dp_overlap_req_start_rows, ) 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 +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager class DpOverlapEagleWithAttProposer(BaseDpOverlapProposer): @@ -61,9 +61,7 @@ def propose_next_overlap( target_next_token_ids1: torch.Tensor, accept_len1: torch.Tensor, draft_step: int, - accept_len_cpu0: torch.Tensor | None = None, - accept_len_cpu1: torch.Tensor | None = None, - accept_len_ready_event: torch.cuda.Event | None = None, + accept_len_cpu: AsyncPinnedCpuTensor | None = None, ) -> EagleSpecProposal: """提交两个 target verify microbatch 的 draft KV,并生成下一轮 proposal。""" @@ -177,20 +175,11 @@ def propose_next_overlap( schedule_scores=schedule_scores, ) - accepted_tail_rows_cpu_by_batch = None if self.backend.is_deepseek_v4 and req_num > 0: - # One shared verify event covers both CPU views of the existing accept-length D2H. - accept_len_ready_event.synchronize() - accepted_tail_rows_cpu_by_batch = [] - for model_input, batch_accept_len_cpu in zip( - model_inputs, - (accept_len_cpu0, accept_len_cpu1), - ): - req_start_rows_cpu = torch.nonzero( - model_input.b_mtp_index_cpu == 0, - as_tuple=False, - ).flatten() - accepted_tail_rows_cpu_by_batch.append(req_start_rows_cpu + batch_accept_len_cpu - 1) + # Both microbatches share one accept-length D2H and its verify event. + accept_len_cpu.wait() + for model_input, batch_accept_len in zip(model_inputs, accept_len_cpu.tensor.split(req_num_by_batch)): + model_input.select_mtp_cpu_mirrors(batch_accept_len) for batch_index, model_input in enumerate(model_inputs): model_input.is_prefill = False @@ -203,15 +192,6 @@ def propose_next_overlap( ) 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 self.backend.is_deepseek_v4 and req_num > 0: - accepted_tail_rows_cpu = accepted_tail_rows_cpu_by_batch[batch_index] - model_input.b_req_idx_cpu = model_input.b_req_idx_cpu.index_select(0, accepted_tail_rows_cpu) - model_input.b_mtp_index_cpu = torch.zeros( - (req_num_by_batch[batch_index],), - dtype=model_input.b_mtp_index_cpu.dtype, - device="cpu", - ) - model_input.b_seq_len_cpu = model_input.b_seq_len_cpu.index_select(0, accepted_tail_rows_cpu) + 1 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 @@ -234,7 +214,7 @@ def propose_next_overlap( draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden draft_seq_lens_by_batch[batch_index].add_(1) if self.backend.is_deepseek_v4 and req_num > 0: - model_inputs[batch_index].b_seq_len_cpu.add_(1) + model_inputs[batch_index].advance_cpu_seq_len() batch_req_num = req_num_by_batch[batch_index] proposal_row_start = proposal_row_offsets[batch_index] diff --git a/lightllm/server/router/model_infer/mtp_speculative/engine.py b/lightllm/server/router/model_infer/mtp_speculative/engine.py index 28a8ab6657..3718a3749e 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/engine.py +++ b/lightllm/server/router/model_infer/mtp_speculative/engine.py @@ -15,6 +15,7 @@ ) from lightllm.server.router.model_infer.mtp_speculative.proposers import build_spec_proposer from lightllm.server.router.model_infer.mtp_speculative.proposers.base import BaseSpecProposer, SpecProposal +from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle_with_att import EagleWithAttProposer if TYPE_CHECKING: from lightllm.server.router.model_infer.mode_backend.base_backend import ModeBackend @@ -107,15 +108,11 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, # [req_num] draft_step: int, accept_len: Optional[torch.Tensor] = None, # [req_num] - accept_len_cpu: Optional[torch.Tensor] = None, # [req_num] - accept_len_ready_event: Optional[torch.cuda.Event] = None, + accept_len_cpu: Optional[AsyncPinnedCpuTensor] = None, # [req_num] ) -> SpecProposal: proposer_kwargs = {} - if self.backend.is_deepseek_v4: - proposer_kwargs = { - "accept_len_cpu": accept_len_cpu, - "accept_len_ready_event": accept_len_ready_event, - } + if self.backend.is_deepseek_v4 and isinstance(self.proposer, EagleWithAttProposer): + proposer_kwargs = {"accept_len_cpu": accept_len_cpu} return self.proposer.propose_next( target_model_input=target_model_input, target_model_output=target_model_output, diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py index 4041a20ec2..0d8c4c0623 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py @@ -53,8 +53,6 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - accept_len_cpu: torch.Tensor | None = None, - accept_len_ready_event: torch.cuda.Event | None = None, ) -> DSparkSpecProposal: """提交 target verify KV,并生成下一轮 DSpark block proposal。 diff --git a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py index 1f554ee3f6..519e6d0ded 100644 --- a/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py +++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/eagle_with_att.py @@ -7,7 +7,7 @@ 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 -from lightllm.server.router.model_infer.pin_mem_manager import g_pin_mem_manager +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager class EagleWithAttProposer(BaseSpecProposer): @@ -43,8 +43,7 @@ def propose_next( b_req_mtp_start_loc: torch.Tensor, draft_step: int, accept_len: torch.Tensor | None = None, - accept_len_cpu: torch.Tensor | None = None, - accept_len_ready_event: torch.cuda.Event | None = None, + accept_len_cpu: AsyncPinnedCpuTensor | None = None, ) -> EagleSpecProposal: """提交验证结果对应的 draft KV,并递归生成下一轮 EAGLE proposal。 @@ -139,19 +138,8 @@ def propose_next( if self.backend.is_deepseek_v4 and req_num > 0: # DSV4's SWA allocator is host-owned. Reuse the accept-length D2H # already issued for post-processing, then keep both mirrors in step. - accept_len_ready_event.synchronize() - req_start_rows_cpu = torch.nonzero( - target_model_input.b_mtp_index_cpu == 0, - as_tuple=False, - ).flatten() - accepted_tail_rows_cpu = req_start_rows_cpu + accept_len_cpu - 1 - draft_input.b_req_idx_cpu = target_model_input.b_req_idx_cpu.index_select(0, accepted_tail_rows_cpu) - draft_input.b_mtp_index_cpu = torch.zeros( - (req_num,), - dtype=target_model_input.b_mtp_index_cpu.dtype, - device="cpu", - ) - draft_input.b_seq_len_cpu = target_model_input.b_seq_len_cpu.index_select(0, accepted_tail_rows_cpu) + 1 + accept_len_cpu.wait() + draft_input.select_mtp_cpu_mirrors(accept_len_cpu.tensor) for step in range(1, draft_step): draft_input.input_ids = draft_token_ids @@ -169,7 +157,7 @@ def propose_next( proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1)) draft_seq_lens.add_(1) if self.backend.is_deepseek_v4 and req_num > 0: - draft_input.b_seq_len_cpu.add_(1) + draft_input.advance_cpu_seq_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 diff --git a/revert.md b/revert.md index 01a16ef7b5..bc67cf8b6f 100644 --- a/revert.md +++ b/revert.md @@ -75,3 +75,7 @@ prompt-logprobs 请求的过滤、命中长度清零和匹配页引用释放统一放在公共加载入口;普通模型和 DSV4 仅在实际缓存加载阶段分派,不再分别实现过滤策略。 普通模型的正常匹配页也在公共入口统一释放,加载子流程保留结束 barrier;DSV4 的正常匹配页继续由异步 session 持有,公共入口只释放其跳过加载的页面。 + +## 2026-09-30 MTP CPU 元数据收敛 + +EAGLE 复用已有 `AsyncPinnedCpuTensor` 传递接受长度和 verify 完成事件,DP overlap 传递一个合并缓冲区,在 proposer 内按 microbatch 请求数拆分视图;不新增 D2H 拷贝或 CUDA event。CPU mirror 的 accepted-tail 选行与序列推进统一由 `ModelInput` 管理,保留 DSV4 host-owned SWA 分配所需的元数据。删除 DSpark 未使用的 CPU 接受长度参数;请求统计仍复用原来的 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 index 0846b7df7f..146b42bd8f 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_dspark.py @@ -2,8 +2,9 @@ import torch +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine 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 +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor, g_pin_mem_manager def test_dspark_prefill_uses_a_shallow_copy_for_target_hidden(): @@ -126,7 +127,10 @@ def forward(model_input): mtp_draft_input_hiddens=None, ) - proposal = proposer.propose_next( + engine = SpecEngine.__new__(SpecEngine) + engine.backend = SimpleNamespace(is_deepseek_v4=True) + engine.proposer = proposer + proposal = engine.propose_next( target_model_input=model_input, target_model_output=SimpleNamespace( mtp_collector=SimpleNamespace(spec_hidden=target_hidden), @@ -135,6 +139,7 @@ def forward(model_input): b_req_mtp_start_loc=torch.tensor([0, 3], dtype=torch.int32), draft_step=2, accept_len=torch.tensor([2, 2], dtype=torch.int32), + accept_len_cpu=AsyncPinnedCpuTensor(torch.tensor([2, 2]), None), ) assert len(forwarded_inputs) == 2 diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py index 9fc9d9d38d..2baf447747 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_overlap.py @@ -3,7 +3,8 @@ import pytest import torch -from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.dp_overlap_engine import DPOverlapSpecEngine 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, @@ -20,6 +21,7 @@ from lightllm.server.router.model_infer.mtp_speculative.proposers.eagle3 import ( Eagle3Proposer, ) +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor class _DraftModel: @@ -101,6 +103,86 @@ def test_dp_overlap_req_start_rows_rejects_nonempty_cpu_input(): ) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("draft_step", [1, 3]) +@pytest.mark.parametrize("empty_batches", [(False, False), (True, False), (False, True), (True, True)]) +def test_dsv4_overlap_eagle_uses_one_cpu_accept_buffer(draft_step, empty_batches): + layouts = ([0, 1, 0, 1, 2], [0, 1]) + accept_lengths = ([2, 1], [1]) + model_inputs, accepts, expected_tails = [], [], [] + for empty, layout, lengths in zip(empty_batches, layouts, accept_lengths): + mtp_index = torch.tensor([] if empty else layout, dtype=torch.int32) + model_input = ModelInput(**vars(_target_input(mtp_index.numel(), mtp_index)), max_q_seq_len=1) + cpu_accept = torch.tensor([] if empty else lengths, dtype=torch.int32) + start_rows = torch.nonzero(mtp_index == 0, as_tuple=False).flatten() + expected_tails.append(model_input.b_seq_len.index_select(0, start_rows + cpu_accept - 1).long()) + model_input.to_cuda() + model_inputs.append(model_input) + accepts.append(cpu_accept.cuda()) + + calls = [] + + def forward(input0, input1): + outputs = [] + for model_input in (input0, input1): + assert torch.equal(model_input.b_req_idx_cpu, model_input.b_req_idx.cpu()) + assert torch.equal(model_input.b_mtp_index_cpu, model_input.b_mtp_index.cpu()) + assert torch.equal(model_input.b_seq_len_cpu, model_input.b_seq_len.cpu()) + outputs.append( + ModelOutput( + logits=model_input.b_seq_len.float().unsqueeze(1), + mtp_collector=ModelMtpOutputCollector( + spec_hidden=torch.ones((model_input.batch_size, 2), device="cuda") + ), + ) + ) + calls.append(True) + return tuple(outputs) + + waits = [] + combined_cpu = torch.zeros(sum(accept.numel() for accept in accepts), dtype=torch.int32) + + def synchronize(): + # Sentinel contents must not be read/split into accepted tails before the wait. + waits.append(True) + combined_cpu.copy_(torch.cat(accepts).cpu()) + + backend = SimpleNamespace( + is_deepseek_v4=True, + draft_models=[SimpleNamespace(_microbatch_overlap_decode_cuda=forward)], + _gen_argmax_token_ids=lambda output: output.logits[:, 0].long(), + ) + engine = DPOverlapSpecEngine.__new__(DPOverlapSpecEngine) + engine.proposer = DpOverlapEagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + outputs = [ + ModelOutput( + logits=torch.empty((model_input.batch_size, 1), device="cuda"), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2), device="cuda")), + ) + for model_input in model_inputs + ] + proposal = engine.propose_next_overlap( + target_model_input0=model_inputs[0], + target_model_output0=outputs[0], + target_next_token_ids0=model_inputs[0].input_ids, + accept_len0=accepts[0], + target_model_input1=model_inputs[1], + target_model_output1=outputs[1], + target_next_token_ids1=model_inputs[1].input_ids, + accept_len1=accepts[1], + draft_step=draft_step, + accept_len_cpu=AsyncPinnedCpuTensor(combined_cpu, SimpleNamespace(synchronize=synchronize)) + if combined_cpu.numel() + else None, + ) + expected = torch.cat(expected_tails)[:, None] + torch.arange(draft_step)[None, :] + torch.testing.assert_close(proposal.token_ids.cpu(), expected) + assert len(calls) == draft_step + assert len(waits) == int(combined_cpu.numel() > 0 and draft_step > 1) + for model_input in model_inputs: + torch.testing.assert_close(model_input.b_seq_len.cpu(), model_input.b_seq_len_cpu) + + def test_overlap_eagle_supports_variable_verify_layout(monkeypatch): _patch_cpu_req_start_rows(monkeypatch) draft_model = _DraftModel() diff --git a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py index 21d61d0f3a..6d243c3327 100644 --- a/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py +++ b/unit_tests/server/router/model_infer/mtp_speculative/test_eagle_with_att.py @@ -3,10 +3,12 @@ import pytest import torch -from lightllm.common.basemodel.batch_objs import ModelMtpOutputCollector, ModelOutput +from lightllm.common.basemodel.batch_objs import ModelInput, ModelMtpOutputCollector, ModelOutput +from lightllm.server.router.model_infer.mtp_speculative.engine import SpecEngine 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 +from lightllm.server.router.model_infer.pin_mem_manager import AsyncPinnedCpuTensor def test_eagle3_reuses_attention_flow_and_maps_proposal_tokens(): @@ -56,6 +58,91 @@ def forward(model_input): assert target_input.mtp_draft_input_hiddens is None +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("draft_step", [1, 2, 3]) +@pytest.mark.parametrize("empty", [False, True]) +def test_dsv4_eagle_reuses_async_accept_lengths_and_preserves_target_mirrors(draft_step, empty): + req_idx = torch.tensor([] if empty else [7, 7, 7, 9, 9], dtype=torch.int32) + mtp_index = torch.tensor([] if empty else [0, 1, 2, 0, 1], dtype=torch.int32) + seq_len = torch.tensor([] if empty else [10, 11, 12, 20, 21], dtype=torch.int32) + target_input = ModelInput( + batch_size=req_idx.numel(), + total_token_num=seq_len.sum().item(), + max_q_seq_len=1, + max_kv_seq_len=21, + input_ids=torch.arange(req_idx.numel(), dtype=torch.int64), + b_req_idx=req_idx, + b_mtp_index=mtp_index, + b_seq_len=seq_len, + b_position_delta=torch.zeros_like(req_idx), + b_shared_seq_len=torch.zeros_like(req_idx), + b_shared_radix_node_id=req_idx.long(), + multimodal_params=[{"images": [], "audios": []} for _ in range(req_idx.numel())], + ) + target_input.to_cuda() + hidden = torch.ones((req_idx.numel(), 2), device="cuda") + accept_len = torch.tensor([] if empty else [2, 1], dtype=torch.int32, device="cuda") + accept_len_cpu = torch.zeros(accept_len.shape, dtype=accept_len.dtype, pin_memory=True) + copy_stream = torch.cuda.Stream() + copy_stream.wait_stream(torch.cuda.current_stream()) + ready_event = torch.cuda.Event() + with torch.cuda.stream(copy_stream): + torch.cuda._sleep(20_000_000) + accept_len_cpu.copy_(accept_len, non_blocking=True) + ready_event.record() + waits = [] + + def synchronize(): + waits.append(True) + ready_event.synchronize() + + calls = [] + + def forward(model_input): + if calls and not empty: + assert torch.equal(model_input.b_req_idx_cpu, model_input.b_req_idx.cpu()) + assert torch.equal(model_input.b_mtp_index_cpu, model_input.b_mtp_index.cpu()) + assert torch.equal(model_input.b_seq_len_cpu, model_input.b_seq_len.cpu()) + calls.append(model_input.b_seq_len.clone()) + return ModelOutput( + logits=model_input.b_seq_len.float().unsqueeze(1), + mtp_collector=ModelMtpOutputCollector(spec_hidden=torch.ones((model_input.batch_size, 2), device="cuda")), + ) + + backend = SimpleNamespace( + is_deepseek_v4=True, + draft_models=[SimpleNamespace(forward=forward)], + _gen_argmax_token_ids=lambda output: output.logits[:, 0].long(), + ) + engine = SpecEngine.__new__(SpecEngine) + engine.backend = backend + engine.proposer = EagleWithAttProposer(backend=backend, enable_dynmaic_mtp=False) + proposal = engine.propose_next( + target_model_input=target_input, + target_model_output=ModelOutput(logits=hidden, mtp_collector=ModelMtpOutputCollector(spec_hidden=hidden)), + target_next_token_ids=target_input.input_ids, + b_req_mtp_start_loc=torch.tensor([] if empty else [0, 3], dtype=torch.int32, device="cuda"), + draft_step=draft_step, + accept_len=accept_len, + accept_len_cpu=AsyncPinnedCpuTensor(accept_len_cpu, SimpleNamespace(synchronize=synchronize)) + if not empty + else None, + ) + expected = ( + torch.empty((0, draft_step), dtype=torch.int64, device="cuda") + if empty + else (torch.tensor([11, 20], device="cuda")[:, None] + torch.arange(draft_step, device="cuda")[None, :]) + ) + torch.testing.assert_close(proposal.token_ids, expected) + assert len(calls) == draft_step + assert len(waits) == int(not empty and draft_step > 1) + assert target_input.b_req_idx_cpu is req_idx + assert target_input.b_mtp_index_cpu is mtp_index + assert target_input.b_seq_len_cpu is seq_len + torch.testing.assert_close(target_input.b_seq_len.cpu(), seq_len) + ready_event.synchronize() + + def test_eagle_with_att_rejects_zero_draft_steps(): proposer = EagleWithAttProposer( backend=SimpleNamespace(draft_models=[]), From 85e243f6d7818ca77cad7b38373143d08f0dd462 Mon Sep 17 00:00:00 2001 From: WANDY666 <1060304770@qq.com> Date: Wed, 30 Sep 2026 02:09:17 +0000 Subject: [PATCH 214/214] refactor(startup): simplify DSV4 cache configuration and restore main PD defaults --- lightllm/server/api_cli.py | 11 +- lightllm/server/api_start.py | 58 ++---- lightllm/server/core/objs/start_args_type.py | 4 +- revert.md | 6 + .../common/test_deepseek4_paged_cache.py | 29 +-- unit_tests/server/test_pd_start_args.py | 187 +++++++++++++++++- 6 files changed, 234 insertions(+), 61 deletions(-) diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e1dbf0c11c..445db0371a 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -136,15 +136,15 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: parser.add_argument( "--pd_kv_page_num", type=int, - default=None, - help="pd mode, kv move page_num; defaults to 8 for DeepSeek-V4 and 16 otherwise.", + default=16, + help="pd mode, kv move page_num", ) parser.add_argument( "--pd_kv_page_size", type=int, - default=None, - help="pd mode, kv page size; defaults to 2048 for DeepSeek-V4 and 1024 otherwise.", + default=1024, + help="pd mode, kv page size.", ) parser.add_argument( @@ -874,7 +874,8 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: "--cpu_cache_token_page_size", type=int, default=None, - help="""The token page size of cpu cache. Defaults to 2048 for DeepSeek-V4 and 256 otherwise.""", + help="""The token page size of cpu cache. Hybrid models use their checkpoint interval. + DeepSeek-V4 defaults to 2048 when big-page checkpoints are disabled; non-hybrid models default to 256.""", ) parser.add_argument( "--cache_placement_strategy", diff --git a/lightllm/server/api_start.py b/lightllm/server/api_start.py index 8a4d2c2ca0..1f6cc164a7 100644 --- a/lightllm/server/api_start.py +++ b/lightllm/server/api_start.py @@ -55,16 +55,6 @@ def _launch_subprocesses(args: StartArgs): auto_set_max_req_total_len(args) auto_set_fused_shared_experts(args) set_unique_server_name(args) - model_type = get_model_type(args.model_dir) - if args.pd_kv_page_num is None: - args.pd_kv_page_num = 8 if model_type == "deepseek_v4" else 16 - if args.pd_kv_page_size is None: - args.pd_kv_page_size = 2048 if model_type == "deepseek_v4" else 1024 - if args.enable_cpu_cache and model_type == "deepseek_v4" and args.llm_kv_type in (None, "None"): - args.llm_kv_type = "fp8kv_dsa" - if args.enable_cpu_cache and model_type == "deepseek_v4" and args.cache_placement_strategy == "adaptive": - logger.warning("DeepSeek-V4 CPU cache does not support adaptive placement; using legacy placement") - args.cache_placement_strategy = "legacy" if args.enable_mps: from lightllm.utils.device_utils import enable_mps @@ -74,15 +64,14 @@ def _launch_subprocesses(args: StartArgs): if args.run_mode not in ["normal", "prefill", "decode", "visual_only"]: return + model_type = get_model_type(args.model_dir) if model_type == "deepseek_v4": - if args.page_size != 256 or args.linear_att_hash_page_size != 256: - logger.warning( - "DeepSeek-V4 forces --page_size and --linear_att_hash_page_size to 256 (got %s and %s)", - args.page_size, - args.linear_att_hash_page_size, - ) + if args.page_size != 256: + logger.warning("DeepSeek-V4 forces --page_size to 256 (got %s)", args.page_size) args.page_size = 256 - args.linear_att_hash_page_size = 256 + if args.enable_cpu_cache and args.cache_placement_strategy == "adaptive": + logger.warning("DeepSeek-V4 incremental CPU cache uses legacy placement") + args.cache_placement_strategy = "legacy" # 通过模型的参数判断是否是多模态模型,包含哪几种模态, 并设置是否启动相应得模块 if args.disable_vision is None: @@ -315,25 +304,21 @@ def _launch_subprocesses(args: StartArgs): # 避免请求释放时将不完整的大页 state 写入 radix cache 并触发断言。 args.linear_att_page_block_num = 10000000 - if ( - args.enable_cpu_cache - and is_hybrid_att_model(args.model_dir) - and get_model_type(args.model_dir) != "deepseek_v4" - ): - args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num - logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for hybrid att model") - elif args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4": - big_page_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num - if big_page_tokens <= args.max_req_total_len: - if args.cpu_cache_token_page_size is None: - args.cpu_cache_token_page_size = big_page_tokens - if args.cpu_cache_token_page_size != big_page_tokens: - raise ValueError("DeepSeek-V4 CPU cache pages must match the hybrid big-page checkpoint interval") - elif args.cpu_cache_token_page_size is None: - args.cpu_cache_token_page_size = 2048 - elif args.enable_cpu_cache and args.cpu_cache_token_page_size is None: - args.cpu_cache_token_page_size = 2048 if get_model_type(args.model_dir) == "deepseek_v4" else 256 if args.enable_cpu_cache: + if model_type == "deepseek_v4": + big_page_tokens = args.linear_att_hash_page_size * args.linear_att_page_block_num + if big_page_tokens <= args.max_req_total_len: + if args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = big_page_tokens + if args.cpu_cache_token_page_size != big_page_tokens: + raise ValueError("DeepSeek-V4 CPU cache pages must match the hybrid big-page checkpoint interval") + elif args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = 2048 + elif is_hybrid_att_model(args.model_dir): + args.cpu_cache_token_page_size = args.linear_att_hash_page_size * args.linear_att_page_block_num + logger.info(f"set cpu_cache_token_page_size to {args.cpu_cache_token_page_size} for hybrid att model") + elif args.cpu_cache_token_page_size is None: + args.cpu_cache_token_page_size = 256 assert ( args.cpu_cache_token_page_size % args.page_size == 0 ), "--cpu_cache_token_page_size must be divisible by --page_size" @@ -529,9 +514,6 @@ def pd_master_start(args: StartArgs): if args.run_mode != "pd_master": return - if args.enable_cpu_cache and get_model_type(args.model_dir) == "deepseek_v4": - raise ValueError("DeepSeek-V4 CPU cache does not support pd_master") - auto_set_max_req_total_len(args) auto_set_response_parsers(args) diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 4998a7f4c8..1f61d788a1 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -211,8 +211,8 @@ class StartArgs: 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: Optional[int] = field(default=None) - pd_kv_page_size: Optional[int] = field(default=None) + pd_kv_page_num: int = field(default=16) + pd_kv_page_size: int = field(default=1024) pd_node_id: int = field(default=-1) enable_cpu_cache: bool = field(default=False) cpu_cache_storage_size: float = field(default=2) diff --git a/revert.md b/revert.md index bc67cf8b6f..ce54ab3972 100644 --- a/revert.md +++ b/revert.md @@ -79,3 +79,9 @@ prompt-logprobs 请求的过滤、命中长度清零和匹配页引用释放统 ## 2026-09-30 MTP CPU 元数据收敛 EAGLE 复用已有 `AsyncPinnedCpuTensor` 传递接受长度和 verify 完成事件,DP overlap 传递一个合并缓冲区,在 proposer 内按 microbatch 请求数拆分视图;不新增 D2H 拷贝或 CUDA event。CPU mirror 的 accepted-tail 选行与序列推进统一由 `ModelInput` 管理,保留 DSV4 host-owned SWA 分配所需的元数据。删除 DSpark 未使用的 CPU 接受长度参数;请求统计仍复用原来的 CPU 缓冲区。 + +## 2026-09-30 启动参数架构收敛 + +恢复 main 的 PD 传输默认值 `16/1024`,移除 DSV4 自动改成 `8/2048` 的调参;显式 CLI 配置保持生效。删除 CPU cache 启动时重复设置 `fp8kv_dsa`、不可达的 CPU 页默认分支,以及 pd_master 的 DSV4 专属 CPU cache 禁止条件(master 本身不创建本地 CPU cache)。 + +必要的模型 KV 参数与 CPU 页/checkpoint 对齐逻辑保留在 `api_start.py` 的对应启动阶段,不新增 `config_utils` helper;模型类型只读取一次并复用。保留物理 KV `page_size=256`;不再覆盖合法的 `linear_att_hash_page_size`,由公共整除校验保证其为 256 的倍数。保留 DSV4 增量 CPU 保存暂时所需的 legacy 策略,以及大页启用时 CPU checkpoint 页长必须相等的约束。 diff --git a/unit_tests/common/test_deepseek4_paged_cache.py b/unit_tests/common/test_deepseek4_paged_cache.py index 2161d3f87a..4ed630ecac 100644 --- a/unit_tests/common/test_deepseek4_paged_cache.py +++ b/unit_tests/common/test_deepseek4_paged_cache.py @@ -611,9 +611,10 @@ def build(): requests.free([source, destination], torch.cat([src, dst])) -@pytest.mark.parametrize("hit_len", [2048, 2304]) +@pytest.mark.parametrize("hash_page_size", [256, 512]) +@pytest.mark.parametrize("hit_extra_page", [False, True]) @pytest.mark.parametrize("chunked", [True, False]) -def test_hybrid_radix_hit_fork_pause_abort_and_eviction(cache, monkeypatch, hit_len, chunked): +def test_hybrid_radix_hit_fork_pause_abort_and_eviction(cache, monkeypatch, hash_page_size, hit_extra_page, chunked): from sortedcontainers import SortedDict from lightllm.server.router.model_infer.infer_batch import InferReq, g_infer_context, CacheTier from lightllm.server.router.dynamic_prompt.hybrid_att_radix_cache import HybridAttPagedRadixCache @@ -621,9 +622,13 @@ def test_hybrid_radix_hit_fork_pause_abort_and_eviction(cache, monkeypatch, hit_ manager, requests = cache small = requests.create_small_page_cache_manager(2) - radix = HybridAttPagedRadixCache(manager.size, 0, 256, 8, manager, small) + radix = HybridAttPagedRadixCache(manager.size, 0, hash_page_size, 2048 // hash_page_size, manager, small) args = get_env_start_args() + args.linear_att_hash_page_size = hash_page_size + args.linear_att_page_block_num = 2048 // hash_page_size args.disable_chunked_prefill = not chunked + tail_len = 2048 + hash_page_size + hit_len = tail_len if hit_extra_page else 2048 if not chunked: args.chunked_prefill_size = args.max_req_total_len for name, value in { @@ -642,13 +647,13 @@ def request(req_idx, total_len): req.shared_kv_node = None req.tail_small_page_buffer_id = None req.hybrid_len_to_big_page_id = SortedDict() - req.hybrid_cache_len = (total_len - 1) // 256 * 256 + req.hybrid_cache_len = (total_len - 1) // hash_page_size * hash_page_size req.image_block_spans = [] req.cache_tiers = {CacheTier.GPU} req.sampling_param = SimpleNamespace(disable_prompt_cache=False) req.prompt_selected_logprobs = SimpleNamespace(copy_capture_slots_if_needed=lambda **kwargs: None) tokens = list(range(total_len)) - hashes = compute_token_list_hash(tokens, 256) + hashes = compute_token_list_hash(tokens, hash_page_size) req.shm_req = SimpleNamespace( input_len=total_len, shm_prompt_ids=SimpleNamespace(arr=tokens), @@ -659,10 +664,10 @@ def request(req_idx, total_len): req.get_chuncked_input_token_len = req.get_chuncked_input_token_len_for_hybrid_att return req - source, held = _reserve(manager, requests, 2305) - req = request(source, 2305) + source, held = _reserve(manager, requests, tail_len + 1) + req = request(source, tail_len + 1) req.hold_kv_len = held.numel() - for end in (2048, 2304) if chunked else (2305,): + for end in (2048, tail_len) if chunked else (tail_len + 1,): assert req.get_chuncked_input_token_len() == end requests.prepare_swa(source, req.cur_kv_len, end) for pool in (manager.swa_pool, manager.c4_pool, manager.c4_indexer_pool, manager.c128_pool): @@ -671,14 +676,14 @@ def request(req_idx, total_len): manager.c4_indexer_state_buffer.uniform_() g_infer_context.save_hybrid_state_to_cache(torch.tensor([source], device="cuda"), [req]) req.cur_kv_len = end - requests.prepare_swa(source, req.cur_kv_len, 2305) - req.cur_kv_len = 2305 + requests.prepare_swa(source, req.cur_kv_len, tail_len + 1) + req.cur_kv_len = tail_len + 1 freed = [] g_infer_context.free_a_req_mem(freed, req) manager.free(torch.cat(freed)) requests.free_req(source) assert manager.swa_page_allocator.can_use_mem_size == manager.swa_num_pages - assert radix.get_tree_total_tokens_num() == 2304 + assert radix.get_tree_total_tokens_num() == tail_len forks = [request(requests.alloc(), hit_len + 1) for _ in range(2)] for fork in forks: @@ -686,7 +691,7 @@ def request(req_idx, total_len): assert fork.cur_kv_len == hit_len assert fork.hold_kv_len == hit_len assert fork.shared_kv_node.node_prefix_total_len == 2048 - if hit_len == 2304: + if hit_extra_page: assert requests.req_to_token_indexs[fork.req_idx, 2048].item() != held[2048].item() restored = requests.req_to_token_indexs[fork.req_idx, :hit_len] for pool, ratio in ((manager.c4_pool, 4), (manager.c4_indexer_pool, 4), (manager.c128_pool, 128)): diff --git a/unit_tests/server/test_pd_start_args.py b/unit_tests/server/test_pd_start_args.py index 3e98fa56a4..a78cd5f9f7 100644 --- a/unit_tests/server/test_pd_start_args.py +++ b/unit_tests/server/test_pd_start_args.py @@ -1,11 +1,141 @@ +import json + import pytest from lightllm.server.api_start import _launch_subprocesses from lightllm.server.core.objs.start_args_type import StartArgs +@pytest.fixture +def startup_args(monkeypatch): + for name in ( + "_set_envs_and_config", + "auto_set_max_req_total_len", + "auto_set_fused_shared_experts", + "set_unique_server_name", + ): + monkeypatch.setattr("lightllm.server.api_start." + name, lambda args: None) + + def validation_finished(args): + raise RuntimeError("startup validation complete") + + monkeypatch.setattr("lightllm.server.api_start.auto_set_response_parsers", validation_finished) + return StartArgs( + model_dir="unused", + max_req_total_len=8192, + eos_id=[2], + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + + @pytest.mark.parametrize( - "page_size,hash_page_size,cpu_page_size", [(1, 256, None), (256, 512, None), (256, 256, 4096), (256, 256, None)] + "model_type,block_num,cpu_page_size,expected", + [ + ("deepseek_v4", 8, None, 4096), + ("deepseek_v4", 8, 4096, 4096), + ("deepseek_v4", 10000000, None, 2048), + ("deepseek_v4", 10000000, 4096, 4096), + ("glm5_next", 8, None, 4096), + ("glm5_next", 8, 256, 4096), + ("llama", 8, None, 256), + ("llama", 8, 1024, 1024), + ], +) +def test_cpu_cache_page_defaults_preserve_model_contracts( + startup_args, monkeypatch, model_type, block_num, cpu_page_size, expected +): + monkeypatch.setattr("lightllm.server.api_start.get_model_type", lambda _: model_type) + monkeypatch.setattr( + "lightllm.server.api_start.is_hybrid_att_model", lambda _: model_type in ("deepseek_v4", "glm5_next") + ) + args = startup_args + args.enable_cpu_cache = True + args.linear_att_page_block_num = block_num + args.cpu_cache_token_page_size = cpu_page_size + with pytest.raises(RuntimeError, match="startup validation complete"): + _launch_subprocesses(args) + assert args.cpu_cache_token_page_size == expected + + +@pytest.mark.parametrize("model_type", ["deepseek_v4", "llama"]) +@pytest.mark.parametrize("enable_cpu_cache", [False, True]) +def test_model_kv_defaults_preserve_hash_pages_kv_selection_and_explicit_pd_tuning( + startup_args, monkeypatch, model_type, enable_cpu_cache +): + monkeypatch.setattr("lightllm.server.api_start.get_model_type", lambda _: model_type) + monkeypatch.setattr("lightllm.server.api_start.is_hybrid_att_model", lambda _: model_type == "deepseek_v4") + args = startup_args + args.enable_cpu_cache = enable_cpu_cache + args.pd_kv_page_num, args.pd_kv_page_size = 8, 2048 + with pytest.raises(RuntimeError, match="startup validation complete"): + _launch_subprocesses(args) + assert args.page_size == (256 if model_type == "deepseek_v4" else 1) + assert args.linear_att_hash_page_size == 512 + assert args.llm_kv_type == "None" + assert (args.pd_kv_page_num, args.pd_kv_page_size) == (8, 2048) + assert args.cache_placement_strategy == ( + "legacy" if enable_cpu_cache and model_type == "deepseek_v4" else "adaptive" + ) + + +def test_disabled_cpu_cache_keeps_page_size_unset(startup_args): + with pytest.raises(RuntimeError, match="startup validation complete"): + _launch_subprocesses(startup_args) + assert startup_args.cpu_cache_token_page_size is None + + +def test_dsv4_cpu_cache_factory_resolves_packed_layout_without_overriding_kv_type(startup_args, monkeypatch, tmp_path): + from lightllm.common.kv_cache_mem_manager import DeepseekV4MemoryManager + from lightllm.common.kv_cache_mem_manager.mem_utils import select_mem_manager_class + from lightllm.utils.envs_utils import get_env_start_args, get_added_mtp_kv_layer_num, get_llm_data_type + from lightllm.utils.llm_utils import get_llm_model_class + from lightllm.utils.kv_cache_utils import calcu_cpu_cache_meta + + (tmp_path / "config.json").write_text( + json.dumps( + { + "model_type": "deepseek_v4", + "num_hidden_layers": 2, + "head_dim": 512, + "index_head_dim": 128, + "compress_ratios": [4, 128], + } + ) + ) + args = startup_args + args.model_dir = str(tmp_path) + args.enable_cpu_cache = True + args.data_type = "bf16" + with pytest.raises(RuntimeError, match="startup validation complete"): + _launch_subprocesses(args) + monkeypatch.setenv("LIGHTLLM_START_ARGS", json.dumps(vars(args))) + cached_functions = ( + get_env_start_args, + get_added_mtp_kv_layer_num, + get_llm_data_type, + get_llm_model_class, + select_mem_manager_class, + calcu_cpu_cache_meta, + ) + for function in cached_functions: + function.cache_clear() + try: + assert select_mem_manager_class() is DeepseekV4MemoryManager + meta = calcu_cpu_cache_meta() + assert meta.token_page_size == 2048 + assert meta.page_shape == (meta.head_dim,) + assert meta.calcu_one_page_size() == meta.head_dim + assert get_env_start_args().llm_kv_type == "None" + finally: + for function in cached_functions: + function.cache_clear() + + +@pytest.mark.parametrize( + "page_size,hash_page_size,cpu_page_size", + [(1, 256, None), (256, 512, None), (256, 256, 4096), (256, 256, None), (256, 512, 4096)], ) def test_dsv4_page_and_checkpoint_configuration(monkeypatch, page_size, hash_page_size, cpu_page_size): for name in ( @@ -20,8 +150,10 @@ def test_dsv4_page_and_checkpoint_configuration(monkeypatch, page_size, hash_pag def validation_finished(args): assert args.page_size == 256 - assert args.linear_att_hash_page_size == 256 - assert args.cpu_cache_token_page_size == 2048 + assert args.linear_att_hash_page_size == hash_page_size + assert args.cpu_cache_token_page_size == hash_page_size * 8 + assert args.llm_kv_type == "None" + assert (args.pd_kv_page_num, args.pd_kv_page_size) == (16, 1024) raise RuntimeError("DSV4 page-size validation passed") monkeypatch.setattr("lightllm.server.api_start.auto_set_response_parsers", validation_finished) @@ -38,7 +170,7 @@ def validation_finished(args): disable_audio=True, disable_shm_warning=True, ) - if cpu_page_size is not None: + if cpu_page_size is not None and cpu_page_size != hash_page_size * 8: with pytest.raises(ValueError, match="CPU cache pages must match"): _launch_subprocesses(args) else: @@ -46,6 +178,53 @@ def validation_finished(args): _launch_subprocesses(args) +@pytest.mark.parametrize("hash_page_size", [128, 384]) +def test_dsv4_rejects_hash_pages_not_aligned_to_physical_kv_pages(monkeypatch, hash_page_size): + for name in ( + "_set_envs_and_config", + "auto_set_max_req_total_len", + "auto_set_fused_shared_experts", + "set_unique_server_name", + ): + monkeypatch.setattr("lightllm.server.api_start." + name, lambda args: None) + monkeypatch.setattr("lightllm.server.api_start.get_model_type", lambda _: "deepseek_v4") + monkeypatch.setattr("lightllm.server.api_start.is_hybrid_att_model", lambda _: True) + args = StartArgs( + model_dir="unused", + linear_att_hash_page_size=hash_page_size, + disable_vision=True, + disable_audio=True, + disable_shm_warning=True, + ) + with pytest.raises(ValueError, match="--linear_att_hash_page_size must be divisible by --page_size"): + _launch_subprocesses(args) + + +def test_pd_cli_and_dataclass_keep_main_defaults(): + import argparse + from lightllm.server.api_cli import add_cli_args + + parsed = add_cli_args(argparse.ArgumentParser()).parse_args([]) + args = StartArgs() + assert (parsed.pd_kv_page_num, parsed.pd_kv_page_size) == (16, 1024) + assert (args.pd_kv_page_num, args.pd_kv_page_size) == (16, 1024) + + +def test_pd_master_does_not_apply_model_local_cpu_cache_restrictions(monkeypatch): + from lightllm.server.api_start import pd_master_start + + monkeypatch.setattr("lightllm.server.api_start._set_envs_and_config", lambda args: None) + monkeypatch.setattr("lightllm.server.api_start.set_unique_server_name", lambda args: None) + + def validation_finished(args): + raise RuntimeError("PD master uses no local CPU cache") + + monkeypatch.setattr("lightllm.server.api_start.auto_set_max_req_total_len", validation_finished) + pd_master_args = StartArgs(run_mode="pd_master", model_dir="unused", enable_cpu_cache=True) + with pytest.raises(RuntimeError, match="PD master uses no local CPU cache"): + pd_master_start(pd_master_args) + + def test_pd_kv_page_size_must_be_divisible_by_model_page_size(monkeypatch): monkeypatch.setattr("lightllm.server.api_start._set_envs_and_config", lambda args: None) monkeypatch.setattr("lightllm.server.api_start.auto_set_max_req_total_len", lambda args: None)