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/basemodel/attention/create_utils.py b/lightllm/common/basemodel/attention/create_utils.py
index 344ba08a1c..272c500319 100644
--- a/lightllm/common/basemodel/attention/create_utils.py
+++ b/lightllm/common/basemodel/attention/create_utils.py
@@ -158,24 +158,28 @@ 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)
def get_neo_prefill_att_backend_class(index=0, priority_list: list = ["fa3", "triton"]) -> BaseAttBackend:
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..45d2760406
--- /dev/null
+++ b/lightllm/common/basemodel/attention/nsa/dsv4_fp8_flashmla_sparse.py
@@ -0,0 +1,222 @@
+import dataclasses
+from typing import TYPE_CHECKING
+
+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)
+
+
+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 _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"],
+ 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)
+ 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,
+ q: torch.Tensor,
+ packed_kv: torch.Tensor,
+ mem_manager,
+ 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,
+ 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
+ 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,
+ flashmla_lse_accum,
+ )
+ 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
+
+ def init_state(self):
+ 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] = flashmla.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,
+ )
+ 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,
+ 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)
+ return out
+
+
+@dataclasses.dataclass
+class _DecodeAttState(BaseDecodeAttState):
+ flashmla_sched_meta: dict = 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.flashmla_sched_meta = {ratio: flashmla.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"]],
+ )
+ return real_out.contiguous()
+
+
+DSV4_NSA_BACKENDS = {"fp8kv_dsa": {"flashmla_sparse": DeepseekV4FlashMlaFp8SparseAttBackend}}
diff --git a/lightllm/common/basemodel/basemodel.py b/lightllm/common/basemodel/basemodel.py
index 5a2d5b77d8..30865f38cd 100755
--- a/lightllm/common/basemodel/basemodel.py
+++ b/lightllm/common/basemodel/basemodel.py
@@ -285,7 +285,7 @@ def _init_cudagraph(self):
cuda_graph_grow_step_size = self.mtp_manager.get_decode_batch_alignment(self.is_mtp_draft_model)
self.graph = (
None
- if self.disable_cudagraph
+ if self.args.run_mode == "prefill" or self.disable_cudagraph
else CudaGraph(
batch_step_size_before_split=cuda_graph_grow_step_size,
split_batch_size=self.args.graph_split_batch_size * decode_tokens_per_request,
@@ -305,7 +305,7 @@ def _init_cudagraph(self):
def _init_prefill_cuda_graph(self):
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:
@@ -383,6 +383,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)
@@ -632,6 +633,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()
@@ -700,6 +702,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)
+ 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):
@@ -972,6 +976,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")
@@ -1047,12 +1053,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_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]
- 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/basemodel/batch_objs.py b/lightllm/common/basemodel/batch_objs.py
index d8f2329b24..93ccf9aeb4 100644
--- a/lightllm/common/basemodel/batch_objs.py
+++ b/lightllm/common/basemodel/batch_objs.py
@@ -40,6 +40,11 @@ class ModelInput:
b_position_delta: torch.Tensor = None
b_prefill_start_loc: torch.Tensor = None
multimodal_params: list = None
+ # cpu 变量
+ 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 不使用。
@@ -51,6 +56,34 @@ 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
+
+ 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_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 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()
@@ -77,6 +110,7 @@ def to_cuda(self):
self.input_ids = self.input_ids.cuda(non_blocking=True)
def __post_init__(self):
+ self.capture_cpu_mirrors()
self.check_input()
def check_input(self):
diff --git a/lightllm/common/basemodel/cuda_graph.py b/lightllm/common/basemodel/cuda_graph.py
index ee888c1a00..ebce0dfd9f 100644
--- a/lightllm/common/basemodel/cuda_graph.py
+++ b/lightllm/common/basemodel/cuda_graph.py
@@ -21,6 +21,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.
@@ -129,6 +143,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 self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool):
model_output = decode_func(infer_state)
self.graph[batch_size] = (graph_obj, infer_state, model_output)
@@ -165,6 +181,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 self.torch_memory_saver.cuda_graph(graph_obj, pool=self.mempool):
model_output, model_output1 = decode_func(infer_state, infer_state1)
self.graph[batch_size] = (
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 0605478e35..27bf4572f7 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/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 fd3de5c9b4..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
@@ -74,7 +74,7 @@ def __init__(
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)
@@ -164,6 +164,32 @@ def experts(
clamp_up_add_one=clamp_up_add_one,
)
+ def experts_with_topk(
+ self,
+ input_tensor: torch.Tensor,
+ topk_weights: torch.Tensor,
+ topk_ids: torch.Tensor,
+ is_prefill: Optional[bool] = None,
+ infer_state=None,
+ alpha: Optional[float] = None,
+ limit: Optional[float] = None,
+ clamp_up_add_one: bool = True,
+ ) -> 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)
+ 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,
+ alpha=alpha,
+ limit=limit,
+ clamp_up_add_one=clamp_up_add_one,
+ )
+
def low_latency_dispatch(
self,
hidden_states: torch.Tensor,
@@ -182,6 +208,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,
@@ -217,7 +260,14 @@ 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,
+ 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(
@@ -227,6 +277,9 @@ def masked_group_gemm(
masked_m=masked_m,
dtype=dtype,
expected_m=expected_m,
+ alpha=alpha,
+ limit=limit,
+ clamp_up_add_one=clamp_up_add_one,
)
def prefilled_group_gemm(
@@ -239,6 +292,9 @@ def prefilled_group_gemm(
recv_topk_weights: torch.Tensor,
hidden_dtype=torch.bfloat16,
microbatch_index: int = 0,
+ 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(
@@ -252,6 +308,9 @@ def prefilled_group_gemm(
w2=self.w2,
hidden_dtype=hidden_dtype,
microbatch_index=microbatch_index,
+ 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 09b538a33d..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
@@ -120,6 +120,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"
@@ -134,6 +146,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,
@@ -200,6 +215,9 @@ def masked_group_gemm(
masked_m: torch.Tensor,
dtype: torch.dtype,
expected_m: int,
+ 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
@@ -212,6 +230,9 @@ def masked_group_gemm(
w2_weight,
w2_scale,
expected_m=expected_m,
+ alpha=alpha,
+ limit=limit,
+ clamp_up_add_one=clamp_up_add_one,
)
def prefilled_group_gemm(
@@ -226,6 +247,9 @@ def prefilled_group_gemm(
w2: WeightPack,
hidden_dtype=torch.bfloat16,
microbatch_index: int = 0,
+ 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
@@ -245,6 +269,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,
+ alpha=alpha,
+ limit=limit,
+ clamp_up_add_one=clamp_up_add_one,
)
else:
gather_out = torch.empty(
@@ -260,7 +287,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)
+ 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 f748aea467..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
@@ -90,6 +90,30 @@ 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,
+ alpha: Optional[float] = None,
+ limit: Optional[float] = None,
+ clamp_up_add_one: bool = True,
+ ):
+ return self._fused_experts(
+ input_tensor=input_tensor,
+ w13=w13,
+ w2=w2,
+ topk_weights=topk_weights,
+ topk_ids=topk_ids,
+ is_prefill=is_prefill,
+ alpha=alpha,
+ limit=limit,
+ clamp_up_add_one=clamp_up_add_one,
+ )
+
def __call__(
self,
input_tensor: torch.Tensor,
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/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py b/lightllm/common/basemodel/triton_kernel/fused_moe/grouped_fused_moe_ep.py
index fc868a85d2..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
@@ -554,8 +554,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/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py b/lightllm/common/basemodel/triton_kernel/quantization/fp8act_quant_kernel.py
index 0a68372887..32a1880f74 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,46 @@ 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)
+
+ 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/basemodel/triton_kernel/redundancy_topk_ids_repair.py b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py
index ba48f414db..692332cc04 100644
--- a/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py
+++ b/lightllm/common/basemodel/triton_kernel/redundancy_topk_ids_repair.py
@@ -22,7 +22,7 @@ def _redundancy_topk_ids_repair_kernel(
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
diff --git a/lightllm/common/kv_cache_mem_manager/__init__.py b/lightllm/common/kv_cache_mem_manager/__init__.py
index b3597b4bc8..d7fdf855ae 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
@@ -18,6 +19,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 1331943b1a..10c6c35edb 100644
--- a/lightllm/common/kv_cache_mem_manager/allocator.py
+++ b/lightllm/common/kv_cache_mem_manager/allocator.py
@@ -2,13 +2,13 @@
from lightllm.server.router.dynamic_prompt.shared_arr import SharedInt
from lightllm.utils.dist_utils import get_current_rank_in_node
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
@@ -24,8 +24,10 @@ def __init__(self, size: int) -> None:
self.can_use_mem_size = self.size
rank_in_node = get_current_rank_in_node()
- # 用共享内存进行共享,router 模块读取进行精确的调度估计;基础层会统一添加服务前缀以防止实例冲突。
- self.shared_can_use_token_num = SharedInt(f"mem_manger_can_use_token_num_{rank_in_node}")
+ # SharedInt adds the service prefix; subpools retain separate counters.
+ if shared_name is None:
+ shared_name = f"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/deepseek2_mem_manager.py b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py
index 4be820b753..26b16dc3c9 100644
--- a/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py
+++ b/lightllm/common/kv_cache_mem_manager/deepseek2_mem_manager.py
@@ -41,6 +41,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,
):
@@ -56,7 +58,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,
@@ -65,6 +67,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
new file mode 100644
index 0000000000..5394ac78e7
--- /dev/null
+++ b/lightllm/common/kv_cache_mem_manager/deepseek4_mem_manager.py
@@ -0,0 +1,995 @@
+import torch
+from dataclasses import dataclass
+from typing import List, Optional, Sequence
+from .mem_manager import MemoryManager
+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_env_start_args, get_unique_server_name
+from lightllm.utils.log_utils import init_logger
+
+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 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
+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_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
+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
+DSV4_C128_STATE_RING = 128 # 128 rows/request before MTP padding
+
+
+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 _DeepseekV4CacheLayout:
+ 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_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_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
+
+ @classmethod
+ def _history_layout(
+ cls,
+ compress_rates: Sequence[int],
+ 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 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 cache expects head_dim={DSV4_MLA_HEAD_DIM}, got {head_dim}")
+ if indexer_head_dim != DSV4_INDEXER_HEAD_DIM:
+ raise ValueError(
+ f"DeepSeek-V4 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
+
+ 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_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
+
+ 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_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
+
+ return dict(
+ 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_page=c4_gpu_pages_per_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=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 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,
+ 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_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,
+ 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
+
+ # 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
+
+ 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,
+ )
+
+
+@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_swa_slots: torch.Tensor
+
+
+class PackedPagePool:
+ """fp8_ds_mla 风格的 packed page 存储: 每页前段连续放 token 的 data 字节,页尾放 per-token scale 字节。
+
+ 寻址是纯 token 槽位 (page = slot // page_size),page 只是 scale-tail/对齐的物理打包技巧,
+ 不存在页粒度的分配。``write``/``read`` 是 torch 参考实现(单测 oracle);生产写入走
+ triton packed writer(destindex_copy_kv_flashmla_dsv4 等),kernel 直接消费 ``buffer``。
+ """
+
+ def __init__(
+ self,
+ 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.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, 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")
+ 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.reshape(-1)
+ 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)
+ scale_range = torch.arange(self.scale_bytes_per_token, device=loc.device)
+ 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.reshape(-1)
+ if loc.numel() == 0:
+ 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)
+ 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)
+
+
+class DeepseekV4MemoryManager(MemoryManager):
+ """DeepSeek-V4 KV cache: 窗口 latent(全层) + c4/c128 压缩 latent(压实层) + c4 indexer-K。
+
+ 与兄弟 manager 一致的 token-slot 设计;req 索引的表都在 DeepseekV4ReqManager。
+
+ - ``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 页分配和释放。
+ - 写入走模型专用的 fused norm/RoPE packed writer,显式接收本轮 SWA 槽;
+ torch codecs 保留为 ABI 的可执行规格(单测 oracle)。
+ """
+
+ operator_class = DeepseekV4MemOperator
+
+ 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,
+ size,
+ dtype,
+ head_num,
+ head_dim,
+ layer_num,
+ 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,
+ 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 (
+ 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"
+ 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)
+ 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_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,
+ head_dim=head_dim,
+ indexer_head_dim=indexer_head_dim,
+ )
+
+ # 全局层号 -> 各压缩池内的压实层号(同 qwen3next 的层号压实手法)
+ 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
+
+ super().__init__(size, dtype, head_num, head_dim, layer_num, always_copy, mem_fraction)
+
+ # ------------------------------------------------------------------ sizing
+ @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 get_cell_size(self):
+ # 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):
+ 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()
+ 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()
+ server = get_unique_server_name()
+
+ 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,
+ layer_num=layer_num,
+ data_bytes=DSV4_MLA_DATA_BYTES_PER_TOKEN,
+ scale_bytes=self.mla_scale_bytes,
+ align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES,
+ )
+ # 注意: 该别名是 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}"
+ )
+ self.HOLD_TOKEN_MEMINDEX = 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.c128_pool: Optional[PackedPagePool] = 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
+ if self.n_c4 > 0:
+ self.c4_pool = PackedPagePool(
+ size=self.c4_size,
+ page_size=DSV4_C4_PAGE_SIZE,
+ layer_num=self.n_c4,
+ 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,
+ )
+ # c4 compressor 在途状态(attention + indexer): swa 页派生寻址(翻译③),随 swa 页
+ # 生灭;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)
+ 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):
+ self._init_state_sentinel(buf)
+ if self.n_c128 > 0:
+ self.c128_pool = PackedPagePool(
+ size=self.c128_size,
+ page_size=DSV4_C128_PAGE_SIZE,
+ layer_num=self.n_c128,
+ data_bytes=DSV4_MLA_DATA_BYTES_PER_TOKEN,
+ scale_bytes=self.mla_scale_bytes,
+ align_bytes=DSV4_MLA_PAGE_ALIGN_BYTES,
+ )
+ # 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"
+ )
+ 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
+
+ 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}) "
+ 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}"
+ )
+
+ # ------------------------------------------------------------------ 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]]
+
+ 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,
+ ) -> 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)
+ 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,
+ resume_swa_slots: torch.Tensor,
+ mem_indexes: Optional[torch.Tensor] = None,
+ ) -> DeepseekV4CpuCacheLoadPlan:
+ """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)
+ 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.swa_pool.buffer.device
+ 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
+
+ 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_swa_slots=resume_swa_slots,
+ )
+
+ def __getstate__(self):
+ state = self.__dict__.copy()
+ # Pinned CPU checkpoints are process-local; IPC readers need only GPU storage.
+ state["big_page_buffers"] = None
+ return state
+
+ 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 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.
+ """
+ block_size = int(block_size)
+ assert block_size > 0
+ assert block_size <= DSV4_SWA_PAGE_SIZE
+ assert token_num % block_size == 0
+
+ req_num = token_num // block_size
+ if req_num == 0:
+ return (
+ torch.empty((0,), dtype=torch.int32, device="cpu"),
+ torch.empty((0,), dtype=torch.int32, device=self.swa_pool.buffer.device),
+ )
+
+ device = self.swa_pool.buffer.device
+ pages_cpu = self.swa_page_allocator.alloc(req_num)
+ pages = pages_cpu.to(device, non_blocking=True)
+ # 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 free_all(self):
+ super().free_all()
+ self.swa_page_allocator.free_all()
+ self.big_page_buffers.clear_to_init_state()
+
+ # ------------------------------------------------------------------ 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)
+ 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(
+ 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].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()
+ 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))
+ 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 :].view(dtype=torch.float32)
+ return (k_fp8 * scale).to(self.dtype)
+
+ # ------------------------------------------------------------------ cache write paths
+ def pack_mla_kv_to_cache_fused_norm_rope(
+ self,
+ layer_index: int,
+ swa_slots: torch.Tensor,
+ kv: torch.Tensor,
+ kv_weight: torch.Tensor,
+ eps: float,
+ freqs_cis: torch.Tensor,
+ positions: torch.Tensor,
+ ):
+ """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,
+ )
+
+ 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
+ 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)
+ 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_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"
+ from lightllm.models.deepseek_v4.triton_kernel.destindex_copy_indexer_k_dsv4 import (
+ destindex_copy_indexer_k_dsv4,
+ )
+
+ destindex_copy_indexer_k_dsv4(
+ indexer_k.reshape(-1, self.indexer_head_dim),
+ mem_index.reshape(-1),
+ positions.reshape(-1),
+ self.c4_indexer_pool.get_layer_buffer(self.layer_to_c4_idx[layer_index]),
+ self.c4_indexer_pool.page_size,
+ )
+
+ # ------------------------------------------------------------------ fenced inherited APIs
+ # kv_buffer 是 page 索引的 uint8 buffer,基类按 token 索引读写的接口会静默写坏数据,显式拦截。
+ def get_index_kv_buffer(self, index):
+ 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 cache does not support token-indexed kv_buffer io")
+
+ def alloc_paged_kv_move_buffer(self, page_num, page_size) -> torch.Tensor:
+ 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,
+ 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
+
+ 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(
+ dp_mems[0],
+ 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,
+ 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
+
+ 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/common/kv_cache_mem_manager/mem_manager.py b/lightllm/common/kv_cache_mem_manager/mem_manager.py
index 38c870ce4e..c2668d1c34 100755
--- a/lightllm/common/kv_cache_mem_manager/mem_manager.py
+++ b/lightllm/common/kv_cache_mem_manager/mem_manager.py
@@ -69,6 +69,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 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 (
@@ -97,8 +100,15 @@ 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()
+ fixed_memory_size = self.get_fixed_memory_size()
pd_kv_move_buffer_size = self.get_pd_kv_move_buffer_size()
- available_memory_bytes = available_memory * 1024 ** 3 - 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 {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)
if world_size > 1:
tensor = torch.tensor(self.size, dtype=torch.int64, device=f"cuda:{get_current_device_id()}")
@@ -106,6 +116,7 @@ def profile_size(self, mem_fraction):
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(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"
@@ -137,6 +148,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,
):
@@ -160,7 +173,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,
@@ -169,6 +182,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/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/__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..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
@@ -78,3 +81,154 @@ def copy_kv_to_mem_manager(self, layer_index: int, mem_index: torch.Tensor, kv:
o_rope,
)
return
+
+
+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
+ 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):
+ raise NotImplementedError("DeepSeek-V4 writes packed KV using request-owned SWA slots")
+
+ 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, source_req_meta, 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
+
+ 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/kv_cache_mem_manager/operator/linear_att.py b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py
index e935539e37..31174071c7 100644
--- a/lightllm/common/kv_cache_mem_manager/operator/linear_att.py
+++ b/lightllm/common/kv_cache_mem_manager/operator/linear_att.py
@@ -94,6 +94,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/kv_cache_mem_manager/qwen3next_mem_manager.py b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py
index 1962fa71b0..aabe1fb3e3 100644
--- a/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py
+++ b/lightllm/common/kv_cache_mem_manager/qwen3next_mem_manager.py
@@ -101,6 +101,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,
):
@@ -111,6 +113,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,
)
@@ -119,7 +123,8 @@ def write_mem_to_page_kv_move_buffer(
helper = self.att_state_page_helper
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,
@@ -128,6 +133,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,
):
@@ -138,6 +145,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/quantization/deepgemm.py b/lightllm/common/quantization/deepgemm.py
index 3c3ee30bb0..6cdc3146d7 100644
--- a/lightllm/common/quantization/deepgemm.py
+++ b/lightllm/common/quantization/deepgemm.py
@@ -62,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)
+ weight, scale = weight_quant(weight.cuda(device), self.block_size, use_ue8m0_scales=True)
output.weight.copy_(weight)
output.weight_scale.copy_(scale)
return
@@ -90,6 +90,7 @@ def apply(
column_major_scales=True,
scale_tma_aligned=True,
alloc_func=alloc_func,
+ use_ue8m0_scales=True,
)
if out is None:
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/common/req_manager/__init__.py b/lightllm/common/req_manager/__init__.py
index ca1f64d5b8..7f2e6ac9b5 100644
--- a/lightllm/common/req_manager/__init__.py
+++ b/lightllm/common/req_manager/__init__.py
@@ -1,4 +1,5 @@
from .base import ReqManager
+from .deepseek4 import DeepseekV4ReqManager
from .linear_att import ReqManagerForMamba
from .glm5_next import Glm5NextReqManager
from .hybrid_base import HybridAttentionReqManager
@@ -6,6 +7,7 @@
__all__ = [
"ReqManager",
+ "DeepseekV4ReqManager",
"HybridAttentionReqManager",
"ReqManagerForMamba",
"Glm5NextReqManager",
diff --git a/lightllm/common/req_manager/deepseek4.py b/lightllm/common/req_manager/deepseek4.py
new file mode 100644
index 0000000000..43f462f5e5
--- /dev/null
+++ b/lightllm/common/req_manager/deepseek4.py
@@ -0,0 +1,187 @@
+from typing import Optional
+
+import torch
+
+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_SWA_PAGE_SIZE,
+ DSV4_PROMPT_CACHE_PAGE_SIZE,
+)
+from lightllm.common.state_cache_manager.deepseek4 import DeepseekV4StateCacheManager
+
+
+class DeepseekV4ReqManager(HybridAttentionReqManager):
+ """Own request-private SWA pages and restore aligned continuation checkpoints.
+
+ 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__(
+ self,
+ max_request_num,
+ max_sequence_length,
+ mem_manager: Optional[DeepseekV4MemoryManager] = None,
+ sliding_window=None,
+ ):
+ super().__init__(max_request_num, max_sequence_length, None)
+ self.sliding_window = sliding_window
+ 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",
+ )
+ 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)
+ )
+ released = sum(p < first_retained for p in pages)
+ return max(0, missing - released)
+
+ def prepare_swa(self, req_idx, start, end):
+ if req_idx == self.HOLD_REQUEST_ID:
+ return
+ 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([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:
+ 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 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
+ )
+
+ 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 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:
+ self.mem_manager.swa_page_allocator.free(list(pages.values()))
+ pages.clear()
+ self.req_to_swa_pages[req_idx].fill_(-1)
+
+ def free(self, free_req_indexes, free_token_index):
+ for req_idx in free_req_indexes:
+ self.clear_runtime_state(req_idx)
+ super().free(free_req_indexes, free_token_index)
+
+ def free_req(self, free_req_index):
+ self.clear_runtime_state(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()
diff --git a/lightllm/common/req_manager/glm5_next.py b/lightllm/common/req_manager/glm5_next.py
index b4814afbf7..04bdf9f2f7 100644
--- a/lightllm/common/req_manager/glm5_next.py
+++ b/lightllm/common/req_manager/glm5_next.py
@@ -29,6 +29,6 @@ def init_hybrid_attention_state(self, req):
super().init_hybrid_attention_state(req)
self.req_to_indexer_tail.buffer[:, req.req_idx].zero_()
- def restore_state(self, req, state_cache_manager, buffer_idx):
- super().restore_state(req, state_cache_manager, buffer_idx)
+ def restore_state(self, req, state_cache_manager, buffer_idx, checkpoint_len=None):
+ super().restore_state(req, state_cache_manager, buffer_idx, checkpoint_len)
self.req_to_indexer_tail.buffer[:, req.req_idx].zero_()
diff --git a/lightllm/common/req_manager/hybrid_base.py b/lightllm/common/req_manager/hybrid_base.py
index 948313116d..02b1fc8be6 100644
--- a/lightllm/common/req_manager/hybrid_base.py
+++ b/lightllm/common/req_manager/hybrid_base.py
@@ -1,5 +1,5 @@
from abc import ABC, abstractmethod
-from typing import TYPE_CHECKING, List
+from typing import TYPE_CHECKING, List, Optional
import torch
@@ -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: Optional[int] = None):
"""将指定大页槽位的 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: Optional[int] = None):
"""将 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..42656a18c5 100644
--- a/lightllm/common/req_manager/linear_att.py
+++ b/lightllm/common/req_manager/linear_att.py
@@ -1,4 +1,4 @@
-from typing import TYPE_CHECKING, List
+from typing import TYPE_CHECKING, List, Optional
import torch
@@ -63,7 +63,13 @@ 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: Optional[List[int]] = None,
+ ):
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 +85,13 @@ 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: Optional[int] = None,
+ ):
# 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 +121,13 @@ 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: Optional[int] = None,
+ ):
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 306a20db80..6ee7b9cfec 100644
--- a/lightllm/common/state_cache_manager/__init__.py
+++ b/lightllm/common/state_cache_manager/__init__.py
@@ -1,7 +1,8 @@
from .base import StateCacheManager
+from .deepseek4 import DeepseekV4StateCacheManager
+from .glm5_next import Glm5NextCacheConfig
from .layer_cache import LayerCache
from .linear_att import LinearAttCacheConfig, LinearAttCacheManager
-from .glm5_next import Glm5NextCacheConfig
def get_hybrid_cache_config():
@@ -12,6 +13,10 @@ def get_hybrid_cache_config():
args = get_env_start_args()
model_cfg, _ = PretrainedConfig.get_config_dict(args.model_dir)
model_type = model_cfg["model_type"]
+ if model_type == "deepseek_v4":
+ from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DeepseekV4CpuCacheLayout
+
+ return DeepseekV4CpuCacheLayout.load_from_args()
if model_type in ("glm5_next", "glm5_next_text"):
return Glm5NextCacheConfig.from_model_config(model_cfg, args)
if model_type in ("qwen3_5", "qwen3_5_moe", "qwen3_5_text", "qwen3_5_moe_text"):
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/distributed/communication_op.py b/lightllm/distributed/communication_op.py
index 93c603212d..11573e3daf 100644
--- a/lightllm/distributed/communication_op.py
+++ b/lightllm/distributed/communication_op.py
@@ -149,6 +149,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/builtin.py b/lightllm/models/builtin.py
index 455b192d03..f3264c1b83 100644
--- a/lightllm/models/builtin.py
+++ b/lightllm/models/builtin.py
@@ -17,6 +17,13 @@ def _tarsier_text_model_is(model_type):
ModelRegistry.register("bloom", "lightllm.models.bloom.model:BloomTpPartModel")
ModelRegistry.register(["deepseek_v2", "deepseek_v3"], "lightllm.models.deepseek2.model:Deepseek2TpPartModel")
ModelRegistry.register(["deepseek_v32"], "lightllm.models.deepseek3_2.model:Deepseek3_2TpPartModel")
+ModelRegistry.register("deepseek_v4", "lightllm.models.deepseek_v4.model:DeepseekV4TpPartModel")
+ModelRegistry.register(
+ "deepseek_v4",
+ "lightllm.models.deepseek_v4.model:DeepseekV4VisionTpPartModel",
+ is_multimodal=True,
+ condition=lambda cfg: cfg.get("vision_n_layers", 0) > 0,
+)
ModelRegistry.register("gemma3", "lightllm.models.gemma3.model:Gemma3TpPartModel")
ModelRegistry.register("gemma4", "lightllm.models.gemma4.model:Gemma4TpPartModel", is_multimodal=True)
ModelRegistry.register("gemma", "lightllm.models.gemma_2b.model:Gemma_2bTpPartModel")
@@ -147,6 +154,8 @@ def _tarsier_text_model_is(model_type):
# Draft keys use the draft checkpoint model_type and the speculative mode.
+DraftModelRegistry.register("deepseek_v4", "eagle_with_att", "lightllm.models.deepseek_v4_mtp.model:DeepseekV4MTPModel")
+DraftModelRegistry.register("deepseek_v4", "dspark", "lightllm.models.deepseek_v4_dspark.model:DeepseekV4DSparkModel")
DraftModelRegistry.register(
"deepseek_v3",
("vanilla_with_att", "eagle_with_att"),
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/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/__init__.py b/lightllm/models/deepseek_v4/__init__.py
new file mode 100644
index 0000000000..e69de29bb2
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..3e730c5b2a
--- /dev/null
+++ b/lightllm/models/deepseek_v4/deepseek_v4_visual.py
@@ -0,0 +1,229 @@
+# 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
+from lightllm.server.visualserver import get_vit_attn_backend
+
+
+@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, 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)
+ 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):
+ 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, cu_seqlens: torch.Tensor) -> torch.Tensor:
+ x = x + self.attn(self.norm1(x), cos, sin, cu_seqlens)
+ 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)
+ cu_seqlens = torch.tensor([0, x.shape[0]], dtype=torch.int32, device=x.device)
+ for block in self.blocks:
+ x = block(x, cos, sin, cu_seqlens)
+ 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/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{dsml_token}invoke>'
+tool_calls_template = "<{dsml_token}{tc_block_name}>\n{tool_calls}\n{dsml_token}{tc_block_name}>"
+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}parameter>
+...
+{dsml_token}invoke>
+<{dsml_token}invoke name="$TOOL_NAME2">
+...
+{dsml_token}invoke>
+{dsml_token}tool_calls>
+
+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}{dsml_token}parameter>'
+ 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"{dsml_token}{tool_calls_block_name}>"
+
+ 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"{dsml_token}invoke"]
+ )
+
+ p_tool_name = re.findall(r'^\s*name="(.*?)">\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"{dsml_token}invoke"]
+ )
+ if content != ">\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/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
new file mode 100644
index 0000000000..a7e6be308d
--- /dev/null
+++ b/lightllm/models/deepseek_v4/infer_struct.py
@@ -0,0 +1,127 @@
+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_*). The full rope tables are
+ model constants and live on the model / layer infers, not here."""
+
+ 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
+ # 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
+ self.dsv4_image_left = None
+ self.dsv4_image_right = 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
+ # 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 _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)
+ 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)
+ # 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)
+ self._dsv4_token_to_batch_idx = torch.repeat_interleave(
+ 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(),
+ )
+ else:
+ 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_image_visibility,
+ build_swa_index,
+ )
+
+ workspace = model.dsv4_workspace
+ self.dsv4_workspace = workspace
+ 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_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_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,
+ )
+ 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,
+ 128,
+ self.dsv4_c128_indices,
+ self.dsv4_c128_lengths,
+ )
+ # 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/__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/compressor.py b/lightllm/models/deepseek_v4/layer_infer/compressor.py
new file mode 100644
index 0000000000..26d64ee65c
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_infer/compressor.py
@@ -0,0 +1,430 @@
+import torch
+import triton
+import triton.language as tl
+from triton.language.extra import libdevice
+
+from lightllm.common.kv_cache_mem_manager.deepseek4_mem_manager import DSV4_SWA_PAGE_SIZE
+
+
+@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,
+ swa_write_slots,
+ 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,
+ 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) + STATE_RING < SWA_PAGE_SIZE
+ if same_page_next and position + STATE_RING < seq_len:
+ return
+ else:
+ if position + COMPRESS_RATIO < seq_len:
+ return
+
+ if IS_C4:
+ 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)
+ 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
+ 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_mtp_index,
+ b_seq_len,
+ b_ready_cache_len,
+ b_q_start_loc,
+ req_to_swa_pages,
+ req_to_swa_stride0,
+ out_slots,
+ norm_weight,
+ rms_eps,
+ cos_table,
+ cos_stride0,
+ cos_stride1,
+ sin_table,
+ sin_stride0,
+ sin_stride1,
+ 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,
+ STATE_RING: tl.constexpr,
+ ROPE_HEAD_DIM: tl.constexpr,
+ FP8_MAX: 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,
+ IS_IN_INDEXER: 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:
+ 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)
+ # 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:
+ swa_page = tl.load(
+ req_to_swa_pages + req_idx * req_to_swa_stride0 + gather_pos // SWA_PAGE_SIZE,
+ 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)
+ else:
+ state_row = req_idx * STATE_RING + gather_pos % STATE_RING
+ state_valid = cache_pos
+ 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_ptrs = (
+ kv_score + current_idx[:, None] * kv_score_stride0 + (head_offset[:, None] + offs[None, :]) * kv_score_stride1
+ )
+ cur_kv = tl.load(
+ cur_ptrs,
+ mask=current_mask,
+ other=0.0,
+ )
+ cur_score = tl.load(
+ 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_ptrs,
+ mask=state_mask,
+ other=0.0,
+ )
+ state_score = tl.load(
+ state_ptrs + STATE_WIDTH,
+ 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)
+ # 以上得到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))
+ 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 * 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)
+
+ 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).
+ 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)
+ 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.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)
+
+ 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 apply_ape(
+ *,
+ kv_score: torch.Tensor,
+ position_ids: torch.Tensor,
+ ape: torch.Tensor,
+ compress_ratio: int,
+):
+ 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),
+ 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,
+ infer_state,
+ layer_idx: int,
+ norm_weight: torch.Tensor,
+ eps: float,
+ head_dim: int,
+ qk_rope_head_dim: int,
+ compress_ratio: int,
+ cos_table: torch.Tensor,
+ sin_table: torch.Tensor,
+ is_in_indexer: bool = False,
+ 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"
+
+ state_buffer = mem_manager.get_c4_indexer_state_buffer(layer_idx)
+ state_ring = mem_manager.c4_state_ring
+ 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:
+
+ 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:
+
+ 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 = 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)
+ 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_swa_pages = infer_state.req_manager.req_to_swa_pages
+
+ _fused_compress_norm_rope_insert_kernel[(kv_score.shape[0],)](
+ kv_score,
+ kv_score.stride(0),
+ kv_score.stride(1),
+ 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_swa_pages,
+ req_to_swa_pages.stride(0),
+ out_slots,
+ norm_weight,
+ eps,
+ cos_table,
+ cos_table.stride(0),
+ cos_table.stride(1),
+ sin_table,
+ sin_table.stride(0),
+ sin_table.stride(1),
+ 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=infer_state.is_prefill,
+ 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,
+ 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,
+ IS_IN_INDEXER=is_in_indexer,
+ num_warps=4,
+ )
+
+ _save_partial_states_kernel[(kv_score.shape[0],)](
+ kv_score,
+ kv_score.stride(0),
+ kv_score.stride(1),
+ infer_state.position_ids,
+ token_to_batch_idx,
+ infer_state.b_req_idx,
+ infer_state.b_seq_len,
+ infer_state.dsv4_swa_write_slots,
+ state_buffer,
+ STATE_WIDTH=state_width,
+ STATE_LAST_DIM=state_last_dim,
+ COMPRESS_RATIO=compress_ratio,
+ IS_C4=is_c4,
+ IS_PREFILL=infer_state.is_prefill,
+ SWA_PAGE_SIZE=DSV4_SWA_PAGE_SIZE,
+ STATE_RING=state_ring,
+ BLOCK=block_state,
+ num_warps=4,
+ )
+ return out_buffer
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..c2592c4d0b
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_infer/hyper_connection.py
@@ -0,0 +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
+
+
+# vllm DeepseekV4DecoderLayer.hc_post_alpha
+HC_POST_ALPHA = 2.0
+
+
+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=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_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, rms_eps, hc_eps, alloc_func):
+ """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),
+ hc_fn,
+ hc_scale,
+ hc_base,
+ out,
+ dim,
+ 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
new file mode 100644
index 0000000000..8e529cffeb
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_infer/post_layer_infer.py
@@ -0,0 +1,28 @@
+from lightllm.models.llama.layer_infer.post_layer_infer import LlamaPostLayerInfer
+from .hyper_connection import hc_head, hc_post
+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: 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,
+ layer_weight.hc_head_scale_.weight,
+ layer_weight.hc_head_base_.weight,
+ cfg["hc_mult"],
+ cfg["hidden_size"],
+ cfg["rms_norm_eps"],
+ cfg.get("hc_eps", 1e-6),
+ self.alloc_tensor,
+ )
+ logits = super().token_forward(collapsed, infer_state, layer_weight)
+ return logits
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..4f00af7fc6
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_infer/pre_layer_infer.py
@@ -0,0 +1,29 @@
+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(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):
+ 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)
+
+ def token_forward(self, input_ids, infer_state: DeepseekV4InferStateInfo, 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
new file mode 100644
index 0000000000..2f4e532428
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_infer/transformer_layer_infer.py
@@ -0,0 +1,1047 @@
+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 .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
+from lightllm.models.deepseek2.triton_kernel.rotary_emb import rotary_emb_fwd
+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
+from lightllm.models.deepseek_v4.workspace import C4_LOGITS_ALIGNMENT, C4_PREFILL_LOGITS_BUDGET_BYTES
+
+
+class DeepseekV4TransformerLayerInfer(Deepseek3_2TransformerLayerInfer):
+ def __init__(self, layer_num, network_config):
+ 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.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.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
+ # 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.num_experts_per_tok = network_config["num_experts_per_tok"]
+ self.routed_scaling_factor = network_config["routed_scaling_factor"]
+ 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(
+ 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_
+ )
+ self.dsv4_prefill_aux_stream = None
+
+ # ------------------------------------------------------------------ forward (HC-threaded)
+ 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."""
+ if torch.is_tensor(input_embdings):
+ residual = input_embdings.view(-1, self.hc_mult, self.embed_dim_)
+ 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_,
+ )
+
+ 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,
+ 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_,
+ )
+
+ 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: 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)
+ 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
+ ):
+ 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)
+
+ 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,
+ handle0.num_unaligned_recv_tokens_per_expert,
+ handle0.recv_src_metadata,
+ recv_x0,
+ recv_indices0,
+ recv_weights0,
+ hidden_dtype=x0.dtype,
+ microbatch_index=infer_state.microbatch_index,
+ alpha=1.0,
+ limit=self.swiglu_limit,
+ clamp_up_add_one=False,
+ )
+
+ 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,
+ handle1.num_unaligned_recv_tokens_per_expert,
+ handle1.recv_src_metadata,
+ recv_x1,
+ recv_indices1,
+ recv_weights1,
+ hidden_dtype=x1.dtype,
+ microbatch_index=infer_state1.microbatch_index,
+ alpha=1.0,
+ limit=self.swiglu_limit,
+ clamp_up_add_one=False,
+ )
+
+ 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,
+ alpha=1.0,
+ limit=self.swiglu_limit,
+ clamp_up_add_one=False,
+ )
+
+ 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,
+ alpha=1.0,
+ limit=self.swiglu_limit,
+ clamp_up_add_one=False,
+ )
+
+ 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:
+ 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,
+ input: torch.Tensor,
+ infer_state: DeepseekV4InferStateInfo,
+ layer_weight: DeepseekV4TransformerLayerWeight,
+ ):
+ 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]
+
+ 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_)
+
+ 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 完成,
+ infer_state.mem_manager.pack_mla_kv_to_cache_fused_norm_rope(
+ layer_index=self.layer_num_,
+ swa_slots=infer_state.dsv4_swa_write_slots,
+ 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 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_]
+ 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]
+ 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)
+
+ # ------------------------------------------------------------------ attention (prefill)
+ 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 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)
+ 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
+ ):
+ if torch.cuda.is_current_stream_capturing():
+ _q = tensor_to_no_ref_tensor(q)
+ _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__()
+ # 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_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):
+ 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)
+ return o
+
+ 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():
+ main_stream = torch.cuda.current_stream()
+ 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)
+ 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,
+ 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,
+ 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
+ # 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
+
+ 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)
+
+ def _context_attention_kernel(
+ 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(
+ nsa_prefill=True,
+ nsa_prefill_dict={
+ "layer_index": self.layer_num_,
+ "compress_ratio": self.compress_ratio,
+ "head_dim_v": self.v_head_dim,
+ "softmax_scale": self.softmax_scale,
+ "attn_sink": layer_weight.attn_sink_.weight,
+ **meta,
+ },
+ )
+ 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),清零以保持确定性
+ attn_out[-pad_q_len:] = 0
+ return attn_out
+
+ # ------------------------------------------------------------------ attention (decode)
+ def token_attention_forward(
+ self, x, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight
+ ):
+ 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(
+ 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)
+ att_control = AttControl(
+ nsa_decode=True,
+ nsa_decode_dict={
+ "layer_index": self.layer_num_,
+ "compress_ratio": self.compress_ratio,
+ "head_dim_v": self.v_head_dim,
+ "softmax_scale": self.softmax_scale,
+ "attn_sink": layer_weight.attn_sink_.weight,
+ **meta,
+ },
+ )
+ 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,
+ infer_state: DeepseekV4InferStateInfo,
+ layer_weight: DeepseekV4TransformerLayerWeight,
+ ):
+ return layer_weight.experts_.experts_with_topk(
+ input_tensor=x,
+ topk_weights=weights,
+ topk_ids=indices,
+ is_prefill=infer_state.is_prefill,
+ infer_state=infer_state,
+ alpha=1.0,
+ limit=float(self.swiglu_limit),
+ clamp_up_add_one=False,
+ )
+
+ 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,
+ alpha=1.0,
+ limit=self.swiglu_limit,
+ clamp_up_add_one=False,
+ )
+ 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_)
+ 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 输出。
+ # 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
+ out = routed + shared
+ return self._tpsp_reduce(input=out, infer_state=infer_state)
+
+ def _select_experts(
+ self, logits, infer_state: DeepseekV4InferStateInfo, layer_weight: DeepseekV4TransformerLayerWeight
+ ):
+ 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
+ indices_dtype = hash_indices_table.dtype
+ 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)
+ topk_softplus_sqrt(
+ weights,
+ indices,
+ logits,
+ self.routed_scaling_factor,
+ bias,
+ input_tokens,
+ hash_indices_table,
+ bias_vl,
+ image_token_start,
+ )
+ return weights, indices
+
+
+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
+ 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"]
+ self.qk_rope_head_dim = network_config["qk_rope_head_dim"]
+ self.eps = network_config["rms_norm_eps"]
+
+ def compress(
+ self,
+ x: torch.Tensor,
+ infer_state: DeepseekV4InferStateInfo,
+ layer_weight: DeepseekV4TransformerLayerWeight,
+ 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,
+ 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
+ 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
+ apply_ape(
+ kv_score=kv_score,
+ position_ids=infer_state.position_ids,
+ ape=ape,
+ compress_ratio=self.compress_ratio,
+ )
+ 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,
+ device=infer_state.mem_index.device,
+ )
+ return fused_compress_op(
+ kv_score=kv_score,
+ infer_state=infer_state,
+ layer_idx=self.layer_idx_,
+ norm_weight=norm_weight,
+ eps=self.eps,
+ 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,
+ is_in_indexer=self.is_in_indexer,
+ out_buffer=out_buffer,
+ )
+
+
+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
+ 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 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.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
+ 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.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; _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
+ else None
+ )
+
+ def write_indexer_k(
+ self,
+ x,
+ infer_state: DeepseekV4InferStateInfo,
+ layer_weight,
+ 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,
+ layer_weight,
+ 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
+
+ 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_,
+ 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,
+ c4_aux_workspace=None,
+ ):
+ swa_indices = infer_state.dsv4_swa_indices.unsqueeze(1)
+ swa_lengths = infer_state.dsv4_swa_lengths
+ 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,
+ 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
+ )
+ elif self.compress_ratio == 128:
+ 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,
+ "extra_indices": extra_indices,
+ "extra_lengths": extra_lengths,
+ }
+
+ def _indexer_q_weight(
+ 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:
+ # 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.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:
+ raise RuntimeError(
+ f"DeepSeek-V4 indexer expects full-token hidden states, got x={x.shape[0]} q_lora={token_num}"
+ )
+ 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 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(
+ 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, 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."""
+ 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
+
+ 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,
+ 4,
+ slots,
+ lengths,
+ )
+ return slots.unsqueeze(1), lengths
+
+ page_size = mem_manager.c4_indexer_pool.page_size
+
+ 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]
+ 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,
+ c4_len,
+ 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 = 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
+
+ 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:
+ 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(
+ 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
+ )
+ 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, 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 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,
+ 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
+
+ self._c4_score_topk(
+ idx_q_fp8,
+ indexer_k_cache,
+ weights,
+ ctx_lens,
+ row_page_table,
+ metadata,
+ c4_cap,
+ 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
+
+ @staticmethod
+ def _c4_score_topk(
+ idx_q_fp8,
+ indexer_k_cache,
+ weights,
+ ctx_lens,
+ row_page_table,
+ metadata,
+ c4_cap,
+ valid_len,
+ top_slots,
+ page_size,
+ logits_out=None,
+ ):
+ args = (
+ idx_q_fp8.unsqueeze(1),
+ indexer_k_cache,
+ weights,
+ ctx_lens,
+ row_page_table,
+ metadata,
+ 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,
+ row_page_table,
+ top_slots,
+ page_size,
+ )
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..9bfb67214f
--- /dev/null
+++ b/lightllm/models/deepseek_v4/layer_weights/transformer_layer_weight.py
@@ -0,0 +1,353 @@
+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,
+ ROWBMMWeight,
+ RMSNormWeight,
+ ParameterWeight,
+ TpAttSinkWeight,
+ FusedMoeWeight,
+)
+from ..triton_kernel.quant_convert import dequant_fp8_block_to_bf16
+
+
+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 (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):
+ super().__init__(layer_num, data_type, network_config, quant_cfg)
+ return
+
+ def _parse_config(self):
+ cfg = self.network_config_
+ 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
+ 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_}"
+
+ def _init_weight(self):
+ self._init_qkvo()
+ if self.has_compressor:
+ self._init_compressor()
+ if self.has_indexer:
+ self._init_indexer()
+ self._init_moe()
+ self._init_norm()
+ self._init_hyper_connection()
+
+ # ------------------------------------------------------------------ attention
+ def _init_qkvo(self):
+ p = f"{self.prefix}.attn"
+ # 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, 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,
+ 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.get_quant_method("wq_b"),
+ )
+ 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 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. 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.
+ 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,
+ 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],
+ weight_names=f"{p}.wo_b.weight",
+ data_type=self.data_type_,
+ quant_method=self.get_quant_method("wo_b"),
+ )
+
+ # ------------------------------------------------------------------ compressor / indexer
+ 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_gate_ = ROWMMWeight(
+ in_dim=self.hidden,
+ 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,
+ 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"
+ # 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,
+ out_dims=[self.index_n_heads],
+ 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)
+ # 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, 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,
+ 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 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],
+ weight_names=f"{p}.gate.weight",
+ data_type=torch.bfloat16,
+ 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,)
+ )
+ 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]
+ # 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,
+ out_dims=[self.hidden],
+ 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",
+ 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):
+ 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):
+ 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 _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 + "."
+ 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 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)
+ 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:
+ weights[k] = dequant_fp8_block_to_bf16(weights[k], weights[scale_k]).to(self.data_type_)
+ del weights[scale_k]
+ # 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)
+ # 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
new file mode 100644
index 0000000000..a90696b360
--- /dev/null
+++ b/lightllm/models/deepseek_v4/model.py
@@ -0,0 +1,480 @@
+import copy
+import importlib.util
+import json
+import math
+import os
+
+import torch
+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.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,
+)
+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 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.deepseek_v4.layer_infer.hyper_connection import hc_post
+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, get_env_start_args
+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
+
+logger = init_logger(__name__)
+
+
+class DeepseekV4TpPartModel(LlamaTpPartModel):
+ req_manager: DeepseekV4ReqManager
+ mem_manager: DeepseekV4MemoryManager
+ has_vision = False
+
+ 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
+
+ def _init_config(self):
+ super()._init_config()
+ normalize_deepseek_v4_config(self.config)
+ return
+
+ 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_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.req_manager = DeepseekV4ReqManager(
+ self.max_req_num,
+ create_max_seq_len,
+ sliding_window=self.config["sliding_window"],
+ )
+ return
+
+ def _get_compress_rates(self, 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()
+ state_mtp_step = 0 if self.args.run_mode == "prefill" else self.args.mtp_step
+ 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,
+ 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"],
+ 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
+ else self.args.cpu_cache_token_page_size
+ ),
+ mem_fraction=self.mem_fraction,
+ )
+ self.req_manager.bind_mem_manager(self.mem_manager)
+ return
+
+ 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")
+ 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,
+ )
+ 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,
+ )
+ 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:
+ 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):
+ self._init_to_get_rotary()
+ 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(
+ 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
+
+ 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:
+ 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.
+ 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 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()
+ 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)
+
+ def _select_mem_indexes(self, model_input: ModelInput):
+ mem_indexes = super()._select_mem_indexes(model_input)
+ self._prepare_dsv4_slots(model_input)
+ return mem_indexes
+
+ 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 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).
+ 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)))
+ freq_exponents = torch.arange(0, dim, 2, dtype=torch.float32, device="cuda") / dim
+ positions = torch.arange(max_seq, dtype=torch.float32, device="cuda")
+
+ 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 = 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
+ 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
+ # 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
+
+
+class DeepseekV4VisionTpPartModel(DeepseekV4TpPartModel):
+ has_vision = True
+
+
+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, 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
+
+ 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 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 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
+
+ # 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}")
+
+ 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 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)
+
+ # 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:
+ 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):
+ fn["arguments"] = json.dumps(fn["arguments"], ensure_ascii=False)
+
+ 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 True
+ thinking_mode = "thinking" if thinking else "chat"
+ effort = kwargs.get("reasoning_effort")
+ if thinking and effort is None:
+ effort = os.getenv("LIGHTLLM_DSV4_THINKING_EFFORT", "high")
+ 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=effort,
+ )
+
+ if tokenize:
+ return self.tokenizer.encode(prompt, add_special_tokens=False)
+ return prompt
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/build_compress_index_dsv4.py b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py
new file mode 100644
index 0000000000..3d9202e292
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/build_compress_index_dsv4.py
@@ -0,0 +1,71 @@
+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,
+ index_ptr,
+ index_stride0,
+ 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.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:
+ 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,
+ ratio: int,
+ index: torch.Tensor,
+ length: torch.Tensor,
+):
+ """Derive compressed entries from group-end full slots in 256-token pages.
+
+ 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]
+ 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),
+ index,
+ index.stride(0),
+ length,
+ cap,
+ RATIO=ratio,
+ BLOCK_E=BLOCK_E,
+ num_warps=4,
+ )
+ return index, length
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..956c4a2ab5
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/build_dspark_swa_index.py
@@ -0,0 +1,111 @@
+import torch
+import triton
+import triton.language as tl
+
+
+@triton.jit
+def _build_dspark_swa_index_kernel(
+ req_idx_ptr,
+ pos_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_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_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=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,
+ 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_swa_pages: 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_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_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_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/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..9d5ce47050
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/build_swa_index_dsv4.py
@@ -0,0 +1,140 @@
+import torch
+import triton
+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,
+ pos_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,
+ 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 < WIDTH
+ offset = pos - w
+ 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)
+ 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_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,
+):
+ """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)
+ 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_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,
+ 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/cache_staging_io.py b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py
new file mode 100644
index 0000000000..0159d39d51
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/cache_staging_io.py
@@ -0,0 +1,265 @@
+import triton
+import triton.language as tl
+
+
+BYTE_BLOCK = 1024
+STATE_BLOCK = 256
+
+
+@triton.jit
+def _pool_pages_kernel(
+ full_slots,
+ pool,
+ pool_stride0,
+ pool_stride1,
+ paired_pool,
+ paired_pool_stride0,
+ paired_pool_stride1,
+ 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 = 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
+ 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 = 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
+ 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,
+ 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_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,
+ 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_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
new file mode 100644
index 0000000000..c7c5996277
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/cpu_cache_io.py
@@ -0,0 +1,737 @@
+import torch
+import triton
+import triton.language as tl
+
+from .cache_staging_io import BYTE_BLOCK as _BYTE_BLOCK
+
+_HISTORY_BLOCK_SIZE = 256
+_C128_RATIO = 128
+
+
+@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 _pack_gpu_cache_to_staging_kernel(
+ full_slots,
+ 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,
+ req_to_swa_pages,
+ req_to_swa_stride0,
+ source_req_meta,
+ 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,
+ 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)
+ 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 = full_slot // 4
+ 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 = 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)
+ 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)
+ 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
+ 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)
+ 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 = 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)
+
+ 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_cpu_cache_to_gpu_kernel(
+ history_c4_slots,
+ history_c128_slots,
+ resume_swa_slots,
+ 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,
+ 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)
+ 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, 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)
+ 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
+
+ _pack_gpu_cache_to_staging_kernel[(page_num * programs_per_page,)](
+ full_slots,
+ 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,
+ 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),
+ 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,
+ 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,
+ )
+
+
+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_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
+
+ 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)
+ if has_c128:
+ assert load_plan.history_c128_slots.shape == (history_block_num, 2)
+
+ 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,
+ 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,
+ 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,
+ )
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..a2aadcdd72
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/norm_rope.cu
@@ -0,0 +1,623 @@
+// 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) {
+ // PDL may start this grid early, so this must precede its first global load.
+ 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;
+
+ 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 =
+ 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];
+
+ 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;
+
+ 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;
+
+ 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;
+
+ 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;
+
+ 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_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/csrc/topk_transform.cu b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu
new file mode 100644
index 0000000000..e29b57fc95
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/csrc/topk_transform.cu
@@ -0,0 +1,348 @@
+// 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;
+
+ // 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;
+ 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;
+
+ 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/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..baa99bf684
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/destindex_copy_indexer_k_dsv4.py
@@ -0,0 +1,101 @@
+import torch
+
+import triton
+import triton.language as tl
+
+
+@triton.jit
+def _fwd_kernel_destindex_copy_indexer_k_dsv4(
+ K,
+ Mem_index,
+ Positions,
+ O_fp8,
+ O_f32,
+ stride_k_bs,
+ stride_k_d,
+ FP8_MIN: tl.constexpr,
+ FP8_MAX: 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)
+ 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 = full_slot // COMPRESS_RATIO
+
+ 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.maximum(tl.max(tl.abs(vals), axis=0), AMAX_MIN)
+ # per-token plain fp32 scale (not ue8m0), matching DeepseekV4MemoryManager._pack_indexer_k
+ 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
+ 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,
+ MemIndex: torch.Tensor,
+ Positions: 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.
+ MemIndex: [T] int — full-token slots for the current rows.
+ Positions: [T] int — logical token positions; only c4 group-end rows are written.
+ 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).
+
+ Bit-compatible with DeepseekV4MemoryManager._pack_indexer_k + PackedPagePool.write.
+ """
+ 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()
+ 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,
+ MemIndex,
+ Positions,
+ 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,
+ AMAX_MIN=1e-4,
+ HEAD_DIM=head_dim,
+ COMPRESS_RATIO=4,
+ 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..054d4a3ca2
--- /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,
+ AMAX_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.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])
+ 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,
+ AMAX_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/dp_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py
new file mode 100644
index 0000000000..cb35e8ea48
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/dp_cache_io.py
@@ -0,0 +1,326 @@
+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
+# 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, source/destination full-slot pointers,
+# source/destination request IDs, logical end.
+_TASK_META_WIDTH = 6
+
+
+@triton.jit
+def _copy_dsv4_dp_caches_kernel(
+ source_pool_ptrs,
+ task_meta,
+ history_meta,
+ 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_c128_pool,
+ dst_c128_pool_stride0,
+ dst_c128_pool_stride1,
+ dst_req_to_swa,
+ req_swa_stride0,
+ dst_swa_pool,
+ dst_swa_pool_stride0,
+ dst_swa_pool_stride1,
+ dst_c4_state,
+ dst_c4_state_stride0,
+ dst_c4_state_stride1,
+ dst_c4_indexer_state,
+ dst_c4_indexer_state_stride0,
+ dst_c4_indexer_state_stride1,
+ 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,
+ 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,
+ 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_index = pid // history_layer_num
+ layer = pid % history_layer_num
+ 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 + 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)
+
+ if HAS_C4:
+ if layer < c4_layer_num:
+ 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 = 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
+
+ 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 * dst_c4_indexer_pool_stride0
+ + src_page * dst_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:
+ 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 = 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
+ dst_token = dst_pool_slot % c128_pool_page_size
+
+ src_page_ptr = (
+ 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
+ )
+ 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)
+
+ 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)
+ source_ptr_row = source_pool_ptrs + source_manager * source_pool_ptr_count
+ 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:
+ page = task_pid % 2
+ layer = task_pid // 2
+ page_i64 = page.to(tl.int64)
+ layer_i64 = layer.to(tl.int64)
+ 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):
+ 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 = 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 + 5).to(tl.pointer_type(tl.uint8))
+ src_c4_indexer_state = tl.load(source_ptr_row + 6).to(tl.pointer_type(tl.uint8))
+ 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
+
+ src_state_row_ptr = (
+ src_c4_state + layer_i64 * dst_c4_state_stride0 + src_state_row * dst_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 * 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_caches(
+ source_pool_ptrs: torch.Tensor,
+ dst_mem_manager,
+ task_meta: torch.Tensor,
+ history_meta: torch.Tensor,
+) -> 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 = 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
+
+ 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
+ dst_c4_indexer_pool = dst_mem_manager.c4_indexer_pool.buffer if has_c4 else None
+ dst_c128_pool = dst_mem_manager.c128_pool.buffer if has_c128 else None
+ dst_swa_pool = dst_mem_manager.swa_pool.buffer
+ dst_c4_state = dst_mem_manager.c4_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_caches_kernel[(program_num,)](
+ source_pool_ptrs,
+ task_meta,
+ history_meta,
+ 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_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.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),
+ dst_c4_state,
+ dst_c4_state.stride(0) if has_c4 else 0,
+ dst_c4_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,
+ 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,
+ 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,
+ 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/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..7bf6910eef
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/gather_c4_indexer_k_dsv4.py
@@ -0,0 +1,75 @@
+import torch
+import triton
+import triton.language as tl
+
+
+@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,
+ page_table_ptr, # [batch, page_cap] int32
+ page_cap,
+ hold_req_id,
+ RATIO: tl.constexpr,
+ PAGE_SIZE: 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 = 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))
+
+
+@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,
+ out: torch.Tensor = None,
+):
+ """Build the logical-c4-page -> physical-c4-page table expected by DeepGEMM paged logits.
+
+ 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
+ which the current token-slot allocator guarantees.
+ """
+ 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
+ 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,
+ req_to_token_indexs,
+ req_to_token_indexs.stride(0),
+ page_table,
+ page_cap,
+ int(hold_req_id),
+ RATIO=4,
+ PAGE_SIZE=page_size,
+ num_warps=1,
+ )
+ return 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
new file mode 100644
index 0000000000..3984a398ef
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/norm_rope_cuda.py
@@ -0,0 +1,102 @@
+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.Tensor = None,
+ weights_out: torch.Tensor = None,
+):
+ 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,
+ 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/pd_cache_io.py b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py
new file mode 100644
index 0000000000..a025bcf1c1
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/pd_cache_io.py
@@ -0,0 +1,413 @@
+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(
+ req_to_swa_pages,
+ req_to_swa_stride0,
+ start_kv_index,
+ 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_rows,
+ 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)
+ 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
+ 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)
+ 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)
+ 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
+ + section_row * 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
+ + section_row * 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,
+ 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_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 + c4_remainder
+ c4_state_start = max(0, request_kv_len - c4_required_rows)
+ 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 = len(rows)
+ c4_rows = torch.tensor(rows, dtype=torch.int64, device="cuda")
+
+ 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,)](
+ 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),
+ 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_rows,
+ 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,
+ 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_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 - 1) // _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,
+ 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/models/deepseek_v4/triton_kernel/quant_convert.py b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py
new file mode 100644
index 0000000000..47d87d4932
--- /dev/null
+++ b/lightllm/models/deepseek_v4/triton_kernel/quant_convert.py
@@ -0,0 +1,16 @@
+import torch
+
+
+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)
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/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/models/deepseek_v4/workspace.py b/lightllm/models/deepseek_v4/workspace.py
new file mode 100644
index 0000000000..5aab324065
--- /dev/null
+++ b/lightllm/models/deepseek_v4/workspace.py
@@ -0,0 +1,238 @@
+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)
+ self.sliding_window = int(model.config["sliding_window"])
+ args = get_env_start_args()
+ 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 = 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
+ self.microbatch_count = 1 + int(overlap)
+
+ 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)
+ 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:
+ 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"
+ )
+
+ 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:
+ 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")
+
+ @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, 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, 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 (
+ 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],
+ )
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..0367243fbf
--- /dev/null
+++ b/lightllm/models/deepseek_v4_dspark/infer_struct.py
@@ -0,0 +1,41 @@
+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."""
+
+ # 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:
+ # 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_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,
+ 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_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..2739dbd20a
--- /dev/null
+++ b/lightllm/models/deepseek_v4_dspark/layer_infer/post_layer_infer.py
@@ -0,0 +1,84 @@
+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,
+)
+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 PostLayerOutput(logits=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 PostLayerOutput(logits=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..0c075073f4
--- /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_dspark.infer_struct import DeepseekV4DSparkInferStateInfo
+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.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: DeepseekV4DSparkInferStateInfo,
+ layer_weight: DeepseekV4DSparkTransformerLayerWeight,
+ ) -> torch.Tensor:
+ """Write target hidden rows into this stage without running draft attention/FFN."""
+ 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_,
+ 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_,
+ 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..22ae2bb5a4
--- /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..6dca7ebc84
--- /dev/null
+++ b/lightllm/models/deepseek_v4_dspark/model.py
@@ -0,0 +1,275 @@
+import gc
+import os
+from typing import List
+
+import torch
+import torch.nn.functional as F
+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.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,
+)
+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.utils.log_utils import init_logger
+
+
+logger = init_logger(__name__)
+
+
+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"
+ 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()
+ 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"])
+ ]
+ 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
+
+ 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):
+ # 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)
+ scratch_pages = model_input.mtp_draft_swa_pages
+ if padded_input is model_input or scratch_pages is None:
+ return padded_input
+
+ 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)
+ 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 = 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
+ 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:
+ # 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:
+ 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.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):
+ 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:
+ 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"
+ 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]
+ 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"
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..2dc54e71fa
--- /dev/null
+++ b/lightllm/models/deepseek_v4_mtp/model.py
@@ -0,0 +1,145 @@
+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.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,
+ 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)]
+ assert self.layers_infer[0].compress_ratio == 0, "DeepSeek-V4 MTP draft layer must be SWA-only"
+ 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_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 MTP 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.0.")})
+ 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 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)
+ 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_cli.py b/lightllm/server/api_cli.py
index 926f9f2030..445db0371a 100644
--- a/lightllm/server/api_cli.py
+++ b/lightllm/server/api_cli.py
@@ -214,6 +214,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
"qwen",
"deepseekv31",
"deepseekv32",
+ "deepseekv4",
"glm47",
"kimi_k2",
"qwen3_coder",
@@ -227,6 +228,7 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser:
choices=[
"deepseek-r1",
"deepseek-v3",
+ "deepseek-v4",
"glm45",
"gpt-oss",
"kimi",
@@ -871,8 +873,9 @@ def add_cli_args(parser: argparse.ArgumentParser) -> 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. 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_models.py b/lightllm/server/api_models.py
index bfb19ff0eb..c1f1749203 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_dict
MAX_SEED = (1 << 63) - 1
@@ -162,7 +163,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_dict(weight_dir)
cls._loaded_defaults = {
"do_sample": generation_cfg.get("do_sample", True),
"presence_penalty": generation_cfg.get("presence_penalty", 0.0),
@@ -244,7 +245,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_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 f258155b4c..cac8f7c3c6 100644
--- a/lightllm/server/api_openai.py
+++ b/lightllm/server/api_openai.py
@@ -174,6 +174,13 @@ def _is_force_thinking_mode(request: ChatCompletionRequest) -> bool:
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-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 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/server/api_start.py b/lightllm/server/api_start.py
index 4d481af2ca..1f6cc164a7 100644
--- a/lightllm/server/api_start.py
+++ b/lightllm/server/api_start.py
@@ -21,6 +21,7 @@
has_vision_module,
is_hybrid_att_model,
auto_set_max_req_total_len,
+ get_model_type,
auto_set_fused_shared_experts,
auto_set_response_parsers,
get_running_max_req_size_per_dp,
@@ -63,6 +64,15 @@ 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:
+ logger.warning("DeepSeek-V4 forces --page_size to 256 (got %s)", args.page_size)
+ args.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:
if has_vision_module(args.model_dir):
@@ -294,10 +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):
- 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")
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"
diff --git a/lightllm/server/build_prompt.py b/lightllm/server/build_prompt.py
index 6e55ebd68f..0b64c4af7e 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/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/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py
index 13ad6367a5..1f61d788a1 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",
@@ -66,6 +67,7 @@ class StartArgs:
"choices": [
"deepseek-r1",
"deepseek-v3",
+ "deepseek-v4",
"glm45",
"gpt-oss",
"kimi",
@@ -214,7 +216,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=256)
+ cpu_cache_token_page_size: Optional[int] = field(default=None)
cache_placement_strategy: str = field(default="adaptive", metadata={"choices": ["adaptive", "legacy"]})
enable_disk_cache: bool = field(default=False)
disk_cache_storage_size: float = field(default=10)
diff --git a/lightllm/server/function_call_parser.py b/lightllm/server/function_call_parser.py
index f4515b0f12..4e20184940 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>",
]
@@ -1934,6 +1935,15 @@ 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):
+ """DeepSeek-V4 uses the V3.2 DSML payload with a tool_calls wrapper."""
+
+ def __init__(self):
+ super().__init__()
+ self.bot_token = f"<{self.dsml_token}tool_calls>"
+ self.eot_token = f"{self.dsml_token}tool_calls>"
+
+
class FunctionCallParser:
"""
Parser for function/tool calls in model outputs.
@@ -1947,6 +1957,7 @@ class FunctionCallParser:
"deepseekv3": DeepSeekV3Detector,
"deepseekv31": DeepSeekV31Detector,
"deepseekv32": DeepSeekV32Detector,
+ "deepseekv4": DeepSeekV4Detector,
"glm47": Glm47Detector,
"kimi_k2": KimiK2Detector,
"llama3": Llama32Detector,
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 fb9e74c47b..7f5ef60347 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_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/multimodal_params.py b/lightllm/server/multimodal_params.py
index ed3535a69f..c3188f7f30 100644
--- a/lightllm/server/multimodal_params.py
+++ b/lightllm/server/multimodal_params.py
@@ -127,6 +127,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 +220,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/pd_io_struct.py b/lightllm/server/pd_io_struct.py
index f479279a26..7eff53c141 100644
--- a/lightllm/server/pd_io_struct.py
+++ b/lightllm/server/pd_io_struct.py
@@ -134,7 +134,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
@@ -142,6 +142,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
@@ -170,6 +171,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/reasoning_parser.py b/lightllm/server/reasoning_parser.py
index cfd2e5a3b1..e647dc7dd0 100644
--- a/lightllm/server/reasoning_parser.py
+++ b/lightllm/server/reasoning_parser.py
@@ -905,6 +905,7 @@ class ReasoningParser:
DetectorMap: Dict[str, Type[BaseReasoningFormatDetector]] = {
"deepseek-r1": DeepSeekR1Detector,
"deepseek-v3": Qwen3Detector,
+ "deepseek-v4": Qwen3Detector,
"glm45": Qwen3Detector,
"gpt-oss": GptOssDetector,
"kimi": KimiDetector,
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/lightllm/server/router/model_infer/infer_batch.py b/lightllm/server/router/model_infer/infer_batch.py
index 99f58ad67b..352c83fd7e 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 TYPE_CHECKING, List, Dict, Tuple, Optional, Callable, Any, Union
-from lightllm.common.req_manager import ReqManager, HybridAttentionReqManager
+from lightllm.common.req_manager import DeepseekV4ReqManager, ReqManager, HybridAttentionReqManager
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
@@ -46,6 +46,7 @@ class InferenceContext:
overlap_stream: torch.cuda.Stream = None # 一些情况下推理进程进行异步折叠操作的异步流对象。
cpu_kv_cache_stream: torch.cuda.Stream = None # 用 cpu kv cache 操作的 stream
is_hybrid_att_model: bool = False # 使用大小页 checkpoint 的混合 attention 模型。
+ is_deepseek_v4: bool = False
def register(
self,
@@ -70,6 +71,7 @@ def register(
self.vocab_size = vocab_size
self.is_hybrid_att_model = isinstance(self.req_manager, HybridAttentionReqManager)
+ self.is_deepseek_v4 = isinstance(self.req_manager, DeepseekV4ReqManager)
return
@@ -138,6 +140,8 @@ def free_a_req_mem(self, free_token_index: List, req: "InferReq"):
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
@@ -371,6 +375,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 = self.get_can_alloc_dsv4_swa_page_num() if self.is_deepseek_v4 else None
for req in paused_reqs:
# 暂停恢复保持原有的保守语义:只有当前完整序列所需的 KV 页面都有足够空间时才恢复,
@@ -379,6 +384,11 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo
if alloc_token_num > can_alloc_token_num:
break
+ 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:
req._hybrid_match_radix_cache()
else:
@@ -390,6 +400,8 @@ def recover_paused_reqs(self, paused_reqs: List["InferReq"], is_master_in_dp: bo
req.shm_req.is_paused = False
logger.debug(f"infer recover paused req id {req.req_id}")
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
return
def get_can_alloc_token_num(self):
@@ -400,39 +412,48 @@ 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):
+ 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()
@@ -443,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
@@ -593,6 +615,10 @@ def __init__(
self.get_chuncked_input_token_len = self.get_chuncked_input_token_len_for_hybrid_att
self.get_chuncked_input_token_ids = self.get_chuncked_input_token_ids_for_hybrid_att
+ 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._init_all_state()
self.generator = None
@@ -626,6 +652,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, HybridAttPagedTreeNode] = None
self.finish_status = FinishStatus()
@@ -634,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
@@ -666,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)
@@ -680,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
)
@@ -696,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:
# 小页匹配
@@ -712,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位置,同时释放
@@ -734,7 +785,8 @@ def _hybrid_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,
)
# 尾部 KV 换到新 mem 后,同步拷贝已捕获的 top-k prompt logprobs。
self.prompt_selected_logprobs.copy_capture_slots_if_needed(
@@ -745,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
@@ -769,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
@@ -828,40 +883,39 @@ 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]
-
- 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)
+ 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
- if chunked_start < self.hybrid_cache_len < end:
- # hybrid checkpoint 对应需要存储的部分。
- end = self.hybrid_cache_len
+ def get_chuncked_input_token_ids(self):
+ return self.shm_req.shm_prompt_ids.arr[0 : self._get_chunked_input_end()]
- return self.shm_req.shm_prompt_ids.arr[0:end]
+ def get_chuncked_input_token_ids_for_hybrid_att(self):
+ 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):
- 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_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):
@@ -967,6 +1021,42 @@ 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_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
+
+ 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)
+ if swa_page_num == 0 or self.args.disable_chunked_prefill:
+ 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
+ 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(),
+ 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
+
+ def get_dsv4_decode_need_swa_page_num(self) -> int:
+ seq_len = self.get_cur_total_len()
+ if seq_len <= 0:
+ return 0
+
+ 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 650cb8493c..47937e793a 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,7 @@
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 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
@@ -150,6 +151,7 @@ def init_model(self, kvargs):
self.model: TpPartBaseModel = self.model # for easy typing
set_random_seed(2147483647)
self.is_hybrid_att_model = isinstance(self.model.req_manager, HybridAttentionReqManager)
+ self.is_deepseek_v4 = isinstance(self.model.req_manager, DeepseekV4ReqManager)
if self.is_hybrid_att_model:
self.small_page_buffers = self.model.req_manager.create_small_page_cache_manager(
@@ -207,6 +209,8 @@ def init_model(self, kvargs):
# 用于协同读取 ShmObjsIOBuffer 中的请求信息的通信tensor和通信组对象。
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")
# 用于在多节点tp模式下协同读取 ShmObjsIOBuffer 中的请求信息的通信tensor和通信组对象。
@@ -289,6 +293,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):
@@ -695,7 +701,7 @@ 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:
self.multi_level_cache_module.update_cpu_cache_task_states()
if req_ids is None:
@@ -722,6 +728,10 @@ def _get_classed_reqs(
prefill_tokens = 0
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
+ if is_deepseek_v4:
+ can_alloc_dsv4_swa_page_num = g_infer_context.get_can_alloc_dsv4_swa_page_num()
for req_obj in ready_reqs:
@@ -754,16 +764,20 @@ def _get_classed_reqs(
is_decode = False
if is_decode:
- # KV 容量检查使用额外分配量,已有页的剩余容量可以覆盖部分或全部 decode 需求。
_, alloc_token_num = req_obj.decode_need_token_num()
- # page_size 较小时,decode 会频繁触发 KV 内存分配。此处额外预申请不超过 8 个 token,
- # 并将数量向下对齐到 page_size 的整数倍,以减少 alloc 调用次数并保持分页分配约束。
+ # Small pages benefit from reserving a few decode slots per allocation.
if alloc_token_num > 0 and self.args.page_size < 8:
alloc_token_num += 8 // self.args.page_size * self.args.page_size
- if alloc_token_num <= can_alloc_token_num:
+ can_run = alloc_token_num <= can_alloc_token_num
+ if can_run and is_deepseek_v4:
+ 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
else:
if wait_pause_count < pause_max_req_num:
if self.args.run_mode == "decode":
@@ -791,17 +805,21 @@ def _get_classed_reqs(
if req_obj.is_slave_req():
continue
- # 计算预算按本轮实际处理的 token 数累计,KV 预算按需要额外分配的页容量扣减。
- token_num, alloc_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, alloc_token_num = req_obj.prefill_need_token_num(is_chuncked_prefill=is_chuncked_prefill)
if prefill_tokens + token_num > self.batch_max_tokens:
continue
- if alloc_token_num <= can_alloc_token_num:
+ can_run = alloc_token_num <= can_alloc_token_num
+ if can_run and is_deepseek_v4:
+ 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
prefill_reqs.append(req_obj)
can_alloc_token_num -= alloc_token_num
+ if is_deepseek_v4:
+ can_alloc_dsv4_swa_page_num -= swa_page_num
else:
if wait_pause_count < pause_max_req_num:
req_obj.wait_pause = True
@@ -855,10 +873,13 @@ 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 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:
@@ -878,6 +899,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/chunked_prefill/impl.py b/lightllm/server/router/model_infer/mode_backend/chunked_prefill/impl.py
index d0026da9bd..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,6 +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=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/dp_shared_kv_trans.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/dp_shared_kv_trans.py
index dedd3bb4a6..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
@@ -11,7 +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_caches
class DPKVSharedMoudle:
@@ -34,11 +36,36 @@ 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 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.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.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,
+ ]
+ )
+ self.dsv4_source_pool_ptrs = torch.tensor(pointer_rows, dtype=torch.uint64, device="cuda")
+ return
+
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 +86,11 @@ 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 = 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(
self.shared_req_infos.arr[0 : len(reqs), :, self._KV_LEN_INDEX], axis=1, keepdims=False
)
@@ -70,16 +102,29 @@ def build_shared_kv_trans_tasks(
is_current_dp_handle = req_dp_rank == self.dp_rank_in_node
# 计算需要传输的 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
+
+ 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
+ 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
- if is_current_dp_handle and trans_size > 0 and alloc_token_num <= g_infer_context.get_can_alloc_token_num():
+ if (
+ is_current_dp_handle
+ and trans_size > 0
+ and alloc_token_num <= g_infer_context.get_can_alloc_token_num()
+ and can_alloc_dsv4_cache
+ ):
assert req.hold_kv_len == req.cur_kv_len
mem_indexes = self.backend._alloc_req_kv_mem(req, alloc_token_num)
assert mem_indexes is not None
# mem_indexes 只描述需要复制的逻辑 KV;页尾预留槽位已经由
# _alloc_req_kv_mem 写入请求表,但不参与本次跨 DP 传输。
mem_indexes = mem_indexes[:trans_size]
+ if self.backend.is_deepseek_v4:
+ dsv4_swa_capacity -= need_swa_pages
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
@@ -94,35 +139,168 @@ 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. 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_gloo_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:
+ 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]
+ 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:
+ 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
+ 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"]):
# 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
+ prompt_cache_page_size = req_manager.get_prompt_cache_page_size()
+ req_list = []
+ seq_list = []
+ 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_list.append(trans_task.req.req_idx)
+ seq_list.append(end)
+ task_meta_data.extend(
+ [
+ trans_task.max_kv_len_mem_manager_index,
+ 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)
+
+ 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 = []
+ 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, 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)]
+ 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
+ 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 = []
+ 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:
+ 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)
for trans_task in trans_tasks:
trans_task.req.cur_kv_len += len(trans_task.mem_indexes)
@@ -138,3 +316,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/dp_backend/impl.py b/lightllm/server/router/model_infer/mode_backend/dp_backend/impl.py
index 2cb4e89e5b..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
@@ -32,8 +32,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:
@@ -69,10 +73,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
)
@@ -496,6 +502,7 @@ def decode_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[InferReq]):
selected_rows = async_selected_row_mask_cpu.tensor.tolist()
run_reqs = [req for req, selected in zip(run_reqs, selected_rows) if selected]
+ mtp_accept_len_cpu = None
if req_num > 0:
next_token_ids, next_token_logprobs = sample(
model_output.logits,
@@ -552,6 +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=(
+ 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(
@@ -824,6 +834,9 @@ def decode_overlap_mtp(self, event_pack: OverlapEventPack, decode_reqs: List[Inf
target_model_output1=model_output1,
target_next_token_ids1=target_next_token_ids1,
accept_len1=mtp_accept_len1,
+ 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/mode_backend/multi_level_kv_cache.py b/lightllm/server/router/model_infer/mode_backend/multi_level_kv_cache.py
index 12a2e9a048..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
+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:
@@ -242,10 +264,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=disk_offload_enable,
)
+ ready_list = [state is CpuPageAllocState.READY_EXISTING for state in alloc_states]
finally:
self.cpu_cache_client.lock.release()
@@ -326,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:
@@ -361,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:
@@ -369,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/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 786c60b30a..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
@@ -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
@@ -51,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
@@ -65,6 +68,16 @@ 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 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:
+ 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
+
# D 节点的请求未结束时会反复进入该过滤逻辑,最多发送 6 次
# PDAbortReq,既提高 abort 消息被传输层处理的概率,也避免持续重复发送。
pd_abort_req_send_count = getattr(req_obj, "pd_abort_req_send_count", 0)
@@ -81,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():
@@ -117,12 +132,58 @@ 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)
+ // 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
+ ):
+ # 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 _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
+ 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
+ # 借用的 full slots 仍由 radix 持有;下一轮从 0 传输时会覆盖请求表。
+ req_obj.cur_kv_len = 0
+ req_obj.hold_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
@@ -134,12 +195,36 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq):
assert req_obj.hold_kv_len == req_obj.cur_kv_len
need_mem_size = req_obj._kv_cache_alloc_need(input_len)
+ 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
+
+ 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)
+
+ prompt_page = req_manager.get_prompt_cache_page_size()
+ 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 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)
assert mem_indexes is not None
# 传输只覆盖真实 KV;最后一个模型页面中尚未使用的部分继续留在
# req_to_token_indexs 中,供后续 decode 直接切片使用。
mem_indexes = mem_indexes[: input_len - req_obj.cur_kv_len]
+ if is_dsv4_req_manager:
+ req_manager.prepare_pd_decode_cache(
+ req_list=[req_obj.req_idx],
+ seq_list=[input_len],
+ )
+ torch.cuda.current_stream().synchronize()
+
while req_obj.pd_trans_kv_start_index < input_len:
cur_page_size = min(trans_page_size, input_len - req_obj.pd_trans_kv_start_index)
# 生成页面传输任务, 放入kv move manager 的处理队列中
@@ -161,7 +246,7 @@ def _decode_node_gen_trans_tasks(self, req_obj: InferReq):
# 混合注意力模型还需接收请求运行态 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=[],
@@ -186,7 +271,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,
@@ -199,23 +284,24 @@ 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 page_kind == "kv":
- req_idx = None
- elif page_kind == "att_state":
- req_idx = req_obj.req_idx
- else:
+ 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", "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,
@@ -234,7 +320,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 b406405e8a..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
@@ -166,17 +166,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_MEMINDEXES[0]],
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
@@ -294,6 +298,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:
@@ -395,6 +400,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 fa644cfc3c..ca5ab86e8c 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 fe3938dcda..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
@@ -66,13 +66,21 @@ 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")
@@ -92,10 +100,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"
@@ -113,7 +125,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)
logger.info(f"Removed remote agent {peer_name}, cost time: {time.monotonic() - start_time:.6f} s")
except BaseException as e:
logger.error(
@@ -259,9 +272,24 @@ 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
+ 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,
@@ -291,7 +319,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 8641e3b69c..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(
@@ -115,10 +115,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":
@@ -128,16 +132,15 @@ def _create_pd_trans_task(
.cpu()
.tolist()
)
- req_idx = None
elif page_kind == "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,
@@ -156,7 +159,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 534966ebe9..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
@@ -145,11 +145,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_MEMINDEXES[0]],
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
@@ -202,12 +205,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,
)
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 5200db6e16..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,12 +153,16 @@ 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_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,)
assert accept_len0.ndim == 1
assert accept_len1.ndim == 1
+ proposer_kwargs = {}
+ 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,
@@ -166,6 +173,7 @@ def propose_next_overlap(
target_next_token_ids1=target_next_token_ids1,
accept_len1=accept_len1,
draft_step=draft_step,
+ **proposer_kwargs,
)
def update_planner_statics(
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 328a6b47de..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,6 +61,7 @@ def propose_next_overlap(
target_next_token_ids1: torch.Tensor,
accept_len1: torch.Tensor,
draft_step: int,
+ accept_len_cpu: AsyncPinnedCpuTensor | None = None,
) -> EagleSpecProposal:
"""提交两个 target verify microbatch 的 draft KV,并生成下一轮 proposal。"""
@@ -174,6 +175,12 @@ def propose_next_overlap(
schedule_scores=schedule_scores,
)
+ if self.backend.is_deepseek_v4 and req_num > 0:
+ # 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
model_input.batch_size = req_num_by_batch[batch_index]
@@ -206,6 +213,8 @@ def propose_next_overlap(
draft_token_ids_by_batch[batch_index] = draft_token_ids
draft_hiddens_by_batch[batch_index] = draft_output.mtp_collector.spec_hidden
draft_seq_lens_by_batch[batch_index].add_(1)
+ if self.backend.is_deepseek_v4 and req_num > 0:
+ 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 83a73c7d9e..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,7 +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[AsyncPinnedCpuTensor] = None, # [req_num]
) -> SpecProposal:
+ proposer_kwargs = {}
+ 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,
@@ -115,6 +120,7 @@ def propose_next(
b_req_mtp_start_loc=b_req_mtp_start_loc,
draft_step=draft_step,
accept_len=accept_len,
+ **proposer_kwargs,
)
# Planner runtime statistics.
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 1d7ca20395..0d8c4c0623 100644
--- a/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py
+++ b/lightllm/server/router/model_infer/mtp_speculative/proposers/dspark.py
@@ -34,11 +34,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
@@ -84,9 +84,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。block 第一行是
# accepted-tail anchor,其余行使用 mask token,由 parallel backbone
@@ -139,6 +140,7 @@ def propose_next(
.contiguous()
)
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:
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 721af1c27d..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,6 +43,7 @@ def propose_next(
b_req_mtp_start_loc: torch.Tensor,
draft_step: int,
accept_len: torch.Tensor | None = None,
+ accept_len_cpu: AsyncPinnedCpuTensor | None = None,
) -> EagleSpecProposal:
"""提交验证结果对应的 draft KV,并递归生成下一轮 EAGLE proposal。
@@ -134,6 +135,12 @@ def propose_next(
draft_input.b_shared_radix_node_id = selected_rows.b_shared_radix_node_id
draft_input.multimodal_params = [{"images": [], "audios": []} for _ in range(req_num)]
+ 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_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
draft_input.mtp_draft_input_hiddens = draft_hidden
@@ -149,6 +156,8 @@ def propose_next(
draft_token_ids = self._gen_argmax_token_ids(draft_output)
proposal_token_ids_by_step.append(draft_token_ids.unsqueeze(1))
draft_seq_lens.add_(1)
+ if self.backend.is_deepseek_v4 and req_num > 0:
+ 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/lightllm/server/tokenizer.py b/lightllm/server/tokenizer.py
index 5dadaf0b13..dfda2f4635 100644
--- a/lightllm/server/tokenizer.py
+++ b/lightllm/server/tokenizer.py
@@ -60,6 +60,17 @@ def get_tokenizer(
try:
tokenizer = AutoTokenizer.from_pretrained(tokenizer_name, trust_remote_code=trust_remote_code, *args, **kwargs)
+ except ValueError as 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 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.
# you can try pip install protobuf==3.20.0 to try repair
@@ -83,6 +94,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, model_cfg)
if model_cfg["architectures"][0] == "TarsierForConditionalGeneration":
from ..models.tarsier2.model import Tarsier2Tokenizer
diff --git a/lightllm/server/visualserver/model_infer/model_rpc.py b/lightllm/server/visualserver/model_infer/model_rpc.py
index 22d5a43e73..6c6e1079eb 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
@@ -95,6 +96,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 8f76b4a53c..5f24a5599b 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
@@ -8,10 +8,86 @@
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_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)
- return json_obj
+ return normalize_deepseek_v4_config(json_obj)
+
+
+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)
+ 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:
@@ -22,7 +98,7 @@ def get_running_max_req_size_per_dp(args) -> int:
local_dp_size = max(1, args.dp // args.nnodes)
# Cache fetch and beam groups need the full capacity.
requires_global_capacity = args.enable_dp_prompt_cache_fetch or args.diverse_mode
- if local_dp_size > 1 and not requires_global_capacity:
+ if local_dp_size > 1 and not requires_global_capacity and is_hybrid_att_model(args.model_dir):
return (args.running_max_req_size + local_dp_size - 1) // local_dp_size
return args.running_max_req_size
@@ -417,6 +493,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
@@ -488,7 +566,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]:
@@ -542,6 +620,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
@@ -567,6 +649,10 @@ def get_reasoning_parser_for_model(model_path: str) -> Optional[str]:
]:
return "qwen3"
+ # 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"
diff --git a/lightllm/utils/envs_utils.py b/lightllm/utils/envs_utils.py
index 66bed05122..c5b88d5799 100644
--- a/lightllm/utils/envs_utils.py
+++ b/lightllm/utils/envs_utils.py
@@ -244,6 +244,11 @@ def get_cache_placement_gpu_capacity_ratio() -> float:
return ratio
+@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():
"""
@@ -304,6 +309,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/lightllm/utils/kv_cache_utils.py b/lightllm/utils/kv_cache_utils.py
index b3047b488e..b4a28ee48f 100644
--- a/lightllm/utils/kv_cache_utils.py
+++ b/lightllm/utils/kv_cache_utils.py
@@ -16,13 +16,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_hybrid_att_model
+from lightllm.utils.config_utils import (
+ 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 (
MemoryManager,
PPLINT8KVMemoryManager,
PPLINT4KVMemoryManager,
Deepseek2MemoryManager,
+ DeepseekV4MemoryManager,
)
from typing import List, Tuple, Optional
@@ -62,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()
@@ -109,11 +116,25 @@ 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:
+ 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,
+ 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模式的精度问题
if not is_hybrid_model:
# 对于非 hybrid 模型,需要额外增加 mtp 的 kv 层数,
@@ -141,11 +162,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
@@ -153,6 +177,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的长度。
diff --git a/revert.md b/revert.md
new file mode 100644
index 0000000000..ce54ab3972
--- /dev/null
+++ b/revert.md
@@ -0,0 +1,87 @@
+# 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` 只多一个汇总清理提交;本文件记录被排除的原始提交,便于后续追溯。
+
+## 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` |
+| 通用采样默认值优化 | 恢复 `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`) |
+
+清理后,`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 持有,公共入口只释放其跳过加载的页面。
+
+## 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/test/unit/test_deepseek_v4_dspark.py b/test/unit/test_deepseek_v4_dspark.py
new file mode 100644
index 0000000000..dbb8e02cdd
--- /dev/null
+++ b/test/unit/test_deepseek_v4_dspark.py
@@ -0,0 +1,257 @@
+import json
+from types import SimpleNamespace
+
+import pytest
+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_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
+ req.get_cur_total_len = lambda: 10
+
+ normal_need = req.get_dsv4_decode_need_swa_page_num()
+ req.args.mtp_mode = "dspark"
+ dspark_need = req.get_dsv4_decode_need_swa_page_num()
+
+ assert dspark_need == normal_need + 1
+
+
+@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",
+ "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_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_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),
+ 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_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),
+ 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
+
+
+@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
+
+ 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"
+
+ 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")
+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_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
+ 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_swa_pages=req_to_swa_pages,
+ 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_swa_slot=120,
+ )
+
+ expected_indices = torch.tensor(
+ [
+ [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",
+ )
+ 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.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(pages_cpu, torch.tensor([2, 0], dtype=torch.int32))
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..ca58304e2d
--- /dev/null
+++ b/tools/convert_deepseek_v4_mtp_to_bf16.py
@@ -0,0 +1,775 @@
+#!/usr/bin/env python3
+"""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, writes quantized and
+low-precision weight tensors as BF16, and preserves source FP32 state tensors.
+
+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
+
+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_dtype: str
+ 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 | None,
+ fp8_block_size: int,
+ mxfp4_block_size: int,
+) -> Tuple[List[TensorSpec], set[str], int]:
+ selected_weight_map = {
+ name: shard for name, shard in weight_map.items() if prefix is None or name.startswith(prefix)
+ }
+ 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 selected_weight_map.items():
+ names_by_shard.setdefault(shard, []).append(name)
+
+ tensor_meta: Dict[str, Tuple[str, Tuple[int, ...]]] = {}
+ original_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_nbytes += _tensor_nbytes(dtype, shape)
+
+ paired_scales: set[str] = set()
+ specs: List[TensorSpec] = []
+ for name in selected_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
+ 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}")
+ 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)
+ output_dtype = "BF16"
+ conversion = "mxfp4"
+ 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")
+
+ specs.append(
+ TensorSpec(
+ name=name,
+ source_shard=selected_weight_map[name],
+ source_dtype=dtype,
+ source_shape=shape,
+ output_shape=output_shape,
+ 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
+ if orphan_scales:
+ 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]]:
+ 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()
+ 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)
+ 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 _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,
+ shard_plan: Sequence[Sequence[TensorSpec]],
+ *,
+ fp8_block_size: int,
+ 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]:
+ 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"))
+ for shard in source_shards
+ }
+ shard_count = len(shard_plan)
+ 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 "
+ 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}.{os.getpid()}.tmp"
+ 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() != 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}")
+
+
+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")
+ 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")
+
+ _, index = _load_index(source_dir)
+ weight_map: Dict[str, str] = index["weight_map"]
+ selected_prefix = None if args.all_weights else args.prefix
+ specs, paired_scales, original_nbytes = _inspect_specs(
+ source_dir,
+ weight_map,
+ 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_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
+
+ print(f"source: {source_dir}")
+ print(f"output: {output_dir}")
+ 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"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
+
+ 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():
+ 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"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_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"
+ 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)
+
+ 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")
+
+ 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")
+ 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",
+ 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",
+ )
+ 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
+
+
+def main() -> None:
+ convert(build_parser().parse_args())
+
+
+if __name__ == "__main__":
+ main()
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()
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/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)
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..4ed630ecac
--- /dev/null
+++ b/unit_tests/common/test_deepseek4_paged_cache.py
@@ -0,0 +1,819 @@
+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,
+ "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,
+ "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_page_num=64,
+ )
+ 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]))
+ for sequences in ([255, 256, 257, 258], [255, 256, 257, 258], [256, 257, 258, 259]):
+ 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 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 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]))
+ 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
+
+
+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 requests.req_to_swa_pages[: requests.HOLD_REQUEST_ID].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 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
+ 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]))
+ 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), torch.tensor([[req_idx, 2048]], device="cuda"), staging)
+ cpu_page = torch.empty(staging.shape, dtype=torch.uint8, pin_memory=True)
+ cpu_page.copy_(staging)
+ 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)
+ )
+ 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 = 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
+ )
+ 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("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
+ 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("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, 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
+ 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, 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 {
+ "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) // 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, hash_page_size)
+ 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, tail_len + 1)
+ req = request(source, tail_len + 1)
+ req.hold_kv_len = held.numel()
+ 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):
+ 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, 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() == tail_len
+
+ 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_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)):
+ 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 multi_level_kv_cache as module
+
+ manager, requests = cache
+ req_idx = requests.alloc()
+ 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)
+ )
+ 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=[],
+ 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]),
+ 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, "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
+ )
+ 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(
+ "gloo", 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_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]),
+ )
+ 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")
+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/common/test_req_manager_page.py b/unit_tests/common/test_req_manager_page.py
index f45ac26ceb..d233c51488 100644
--- a/unit_tests/common/test_req_manager_page.py
+++ b/unit_tests/common/test_req_manager_page.py
@@ -61,6 +61,7 @@ def _make_context(monkeypatch):
monkeypatch.setattr(base_backend, "g_infer_context", context)
backend = base_backend.ModeBackend.__new__(base_backend.ModeBackend)
backend.args = context.args
+ backend.is_deepseek_v4 = False
return context, backend
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_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..1c1e0f98af
--- /dev/null
+++ b/unit_tests/models/deepseek_v4/test_vision_integration.py
@@ -0,0 +1,468 @@
+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, 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, 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_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),
+ 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._hybrid_match_radix_cache()
+
+ assert radix_cache.calls == [
+ (1024, False),
+ (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, get_swa_page_need=lambda req, start, end: 16
+ ),
+ )
+
+ 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.req_idx = 0
+
+ assert req.get_dsv4_recover_need_swa_page_num() == 8
+
+
+@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.multi_level_kv_cache import (
+ MultiLevelKvCacheModule,
+ )
+
+ req = SimpleNamespace(image_block_spans=spans)
+
+ 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 (
+ 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, resume_swa_slots, mem_indexes):
+ prepare_calls.append((token_num, loaded_end))
+ return SimpleNamespace(mem_indexes=mem_indexes)
+
+ 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,
+ 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),
+ 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,
+ alloc=lambda token_num: torch.arange(token_num, dtype=torch.int32),
+ )
+ 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,
+ 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)],
+ 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]),
+ 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(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,
+ "get_can_alloc_dsv4_swa_page_num",
+ lambda: 2,
+ )
+
+ 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, 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
+ 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_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")
+ image_right = torch.tensor([0, 2, 0], dtype=torch.int32, device="cuda")
+
+ build_swa_index(
+ req_idx,
+ positions,
+ req_to_swa_pages,
+ output,
+ lengths,
+ write_slots,
+ window=4,
+ image_left=image_left,
+ image_right=image_right,
+ )
+
+ assert output.cpu().tolist() == [
+ [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/dynamic_prompt/test_radix_cache.py b/unit_tests/server/router/dynamic_prompt/test_radix_cache.py
index 80265ac40d..bdb877773e 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
@@ -317,5 +319,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()
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 233a7345e0..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
@@ -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)
@@ -292,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 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"]
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 89e15fc93d..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))
@@ -117,6 +119,7 @@ def test_dp_cache_fetch_reserves_pages_and_keeps_logical_transfer_size(monkeypat
)
source_table = torch.arange(32, dtype=torch.int32).reshape(2, 16)
backend = SimpleNamespace(
+ is_deepseek_v4=False,
model=SimpleNamespace(mem_manager=SimpleNamespace()),
dp_world_size=1,
rank_in_dp=0,
diff --git a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py
index 73e133cc18..cb1f6dff86 100644
--- a/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py
+++ b/unit_tests/server/router/model_infer/mode_backend/test_pd_dynamic_split.py
@@ -45,6 +45,7 @@ def _classify_without_token_capacity(monkeypatch, req, support_overlap=True):
)
backend.support_overlap = support_overlap
backend.is_master_in_dp = True
+ backend.is_deepseek_v4 = False
logger = MagicMock()
backend.logger = logger
backend._timer_merge_radix_tree = MagicMock()
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 63e78e4e5d..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():
@@ -36,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)
@@ -96,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),
@@ -105,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 d5dae90b4b..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,10 +103,91 @@ 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()
backend = SimpleNamespace(
+ is_deepseek_v4=False,
max_draft_step=2,
draft_models=[draft_model],
model=SimpleNamespace(
@@ -145,6 +228,7 @@ def test_overlap_eagle_supports_variable_verify_layout(monkeypatch):
def test_overlap_eagle_supports_empty_verify_rows():
draft_model = _DraftModel()
backend = SimpleNamespace(
+ is_deepseek_v4=False,
max_draft_step=2,
draft_models=[draft_model],
model=SimpleNamespace(
@@ -182,6 +266,7 @@ def test_overlap_eagle_returns_dynamic_schedule_scores(monkeypatch):
_patch_cpu_req_start_rows(monkeypatch)
draft_model = _DraftModel()
backend = SimpleNamespace(
+ is_deepseek_v4=False,
max_draft_step=2,
draft_models=[draft_model],
model=SimpleNamespace(
@@ -232,6 +317,7 @@ def test_overlap_eagle_no_att_supports_dynamic_draft_step():
device = "cuda"
draft_model = _DraftModel()
backend = SimpleNamespace(
+ is_deepseek_v4=False,
max_draft_step=3,
draft_models=[draft_model],
_gen_argmax_token_ids_and_prob=lambda output: (
@@ -284,6 +370,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch):
_patch_cpu_req_start_rows(monkeypatch)
draft_model = _DraftModel()
backend = SimpleNamespace(
+ is_deepseek_v4=False,
max_draft_step=2,
draft_models=[draft_model],
model=SimpleNamespace(
@@ -328,6 +415,7 @@ def test_autoregressive_eagle_reuses_overlap_inputs(monkeypatch):
def test_eagle3_maps_draft_token_ids_in_proposer():
proposer = Eagle3Proposer.__new__(Eagle3Proposer)
proposer.backend = SimpleNamespace(
+ is_deepseek_v4=False,
draft_models=[SimpleNamespace(map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100)],
_gen_argmax_token_ids=lambda _: torch.tensor([1, 2]),
_gen_argmax_token_ids_and_prob=lambda _: (
@@ -347,6 +435,7 @@ def test_eagle3_maps_draft_token_ids_in_proposer():
def test_dp_overlap_eagle3_maps_draft_token_ids_in_proposer():
proposer = DpOverlapEagle3Proposer.__new__(DpOverlapEagle3Proposer)
proposer.backend = SimpleNamespace(
+ is_deepseek_v4=False,
draft_models=[SimpleNamespace(map_draft_vocab_to_main_vocab=lambda token_ids: token_ids + 100)],
_gen_argmax_token_ids=lambda _: torch.tensor([1, 2]),
_gen_argmax_token_ids_and_prob=lambda _: (
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 365acb30b2..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=[]),
@@ -158,6 +245,7 @@ def forward(model_input):
)
backend = SimpleNamespace(
+ is_deepseek_v4=False,
draft_models=[SimpleNamespace(forward=forward)],
_gen_argmax_token_ids=lambda output: output.logits[:, 0].long(),
_gen_argmax_token_ids_and_prob=lambda output: (
diff --git a/unit_tests/server/test_mtp_start_args.py b/unit_tests/server/test_mtp_start_args.py
index 87e12479d4..0c279c7ef8 100644
--- a/unit_tests/server/test_mtp_start_args.py
+++ b/unit_tests/server/test_mtp_start_args.py
@@ -4,6 +4,10 @@
from lightllm.server.core.objs.start_args_type import StartArgs
+class _StartupStopped(Exception):
+ pass
+
+
def test_mtp_requires_cuda_graph(monkeypatch):
monkeypatch.setattr("lightllm.server.api_start._set_envs_and_config", lambda args: None)
args = StartArgs(mtp_mode="vanilla_no_att", disable_cudagraph=True)
@@ -30,3 +34,21 @@ def test_mtp_prefill_still_requires_positive_step(monkeypatch):
with pytest.raises(AssertionError):
_launch_subprocesses(args)
+
+
+def test_dsv4_cpu_cache_falls_back_to_legacy_placement(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)
+ monkeypatch.setattr("lightllm.server.api_start.auto_set_fused_shared_experts", lambda args: None)
+ monkeypatch.setattr("lightllm.server.api_start.set_unique_server_name", lambda args: None)
+ monkeypatch.setattr("lightllm.server.api_start.get_model_type", lambda model_dir: "deepseek_v4")
+ monkeypatch.setattr(
+ "lightllm.server.api_start.has_vision_module",
+ lambda model_dir: (_ for _ in ()).throw(_StartupStopped),
+ )
+ args = StartArgs(enable_cpu_cache=True, disable_audio=True)
+
+ with pytest.raises(_StartupStopped):
+ _launch_subprocesses(args)
+
+ assert args.cache_placement_strategy == "legacy"
diff --git a/unit_tests/server/test_pd_start_args.py b/unit_tests/server/test_pd_start_args.py
index def9608eb8..a78cd5f9f7 100644
--- a/unit_tests/server/test_pd_start_args.py
+++ b/unit_tests/server/test_pd_start_args.py
@@ -1,9 +1,230 @@
+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(
+ "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 (
+ "_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.page_size == 256
+ 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)
+ 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 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:
+ with pytest.raises(RuntimeError, match="DSV4 page-size validation passed"):
+ _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)