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' +tool_calls_template = "<{dsml_token}{tc_block_name}>\n{tool_calls}\n" +tool_calls_block_name: str = "tool_calls" + +tool_output_template: str = "{content}" + +REASONING_EFFORT_MAX = ( + "Reasoning Effort: Absolute maximum with no shortcuts permitted.\n" + "You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n" # noqa: E501 + "Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n" # noqa: E501 +) + +TOOLS_TEMPLATE = """## Tools + +You have access to a set of tools to help answer the user's question. You can invoke tools by writing a \ +"<{dsml_token}tool_calls>" block like the following: + +<{dsml_token}tool_calls> +<{dsml_token}invoke name="$TOOL_NAME"> +<{dsml_token}parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE +... + +<{dsml_token}invoke name="$TOOL_NAME2"> +... + + + +String parameters should be specified as is and set `string="true"`. For all other types (numbers, \ +booleans, arrays, objects), pass the value in JSON format and set `string="false"`. + +If thinking_mode is enabled (triggered by {thinking_start_token}), you MUST output your complete \ +reasoning inside {thinking_start_token}...{thinking_end_token} BEFORE any tool calls or final response. + +Otherwise, output directly after {thinking_end_token} with tool calls or final response. + +### Available Tool Schemas + +{tool_schemas} + +You MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. +""" + +# ============================================================ +# Utility Functions +# ============================================================ + + +def to_json(value: Any) -> str: + """Serialize a value to JSON string.""" + try: + return json.dumps(value, ensure_ascii=False) + except: + return json.dumps(value, ensure_ascii=True) + + +def tools_from_openai_format(tools): + """Extract function definitions from OpenAI-format tool list.""" + return [tool["function"] for tool in tools] + + +def tool_calls_from_openai_format(tool_calls): + """Convert OpenAI-format tool calls to internal format.""" + return [ + { + "name": tool_call["function"]["name"], + "arguments": tool_call["function"]["arguments"], + } + for tool_call in tool_calls + ] + + +def tool_calls_to_openai_format(tool_calls): + """Convert internal tool calls to OpenAI format.""" + return [ + { + "type": "function", + "function": { + "name": tool_call["name"], + "arguments": tool_call["arguments"], + }, + } + for tool_call in tool_calls + ] + + +def encode_arguments_to_dsml(tool_call: Dict[str, str]) -> str: + """ + Encode tool call arguments into DSML parameter format. + + Args: + tool_call: Dict with "name" and "arguments" (JSON string) keys. + + Returns: + DSML-formatted parameter string. + """ + p_dsml_template = '<{dsml_token}parameter name="{key}" string="{is_str}">{value}' + P_dsml_strs = [] + + try: + arguments = json.loads(tool_call["arguments"]) + except Exception: + arguments = {"arguments": tool_call["arguments"]} + + for k, v in arguments.items(): + p_dsml_str = p_dsml_template.format( + dsml_token=dsml_token, + key=k, + is_str="true" if isinstance(v, str) else "false", + value=v if isinstance(v, str) else to_json(v), + ) + P_dsml_strs.append(p_dsml_str) + + return "\n".join(P_dsml_strs) + + +def decode_dsml_to_arguments(tool_name: str, tool_args: Dict[str, Tuple[str, str]]) -> Dict[str, str]: + """ + Decode DSML parameters back to a tool call dict. + + Args: + tool_name: Name of the tool. + tool_args: Dict mapping param_name -> (value, is_string_flag). + + Returns: + Dict with "name" and "arguments" (JSON string) keys. + """ + + def _decode_value(key: str, value: str, string: str): + if string == "true": + value = to_json(value) + return f"{to_json(key)}: {value}" + + tool_args_json = "{" + ", ".join([_decode_value(k, v, string=is_str) for k, (v, is_str) in tool_args.items()]) + "}" + return dict(name=tool_name, arguments=tool_args_json) + + +def render_tools(tools: List[Dict[str, Union[str, Dict[str, Any]]]]) -> str: + """ + Render tool schemas into the system prompt format. + + Args: + tools: List of tool schema dicts (each with name, description, parameters). + + Returns: + Formatted tools section string. + """ + tools_json = [to_json(t) for t in tools] + + return TOOLS_TEMPLATE.format( + tool_schemas="\n".join(tools_json), + dsml_token=dsml_token, + thinking_start_token=thinking_start_token, + thinking_end_token=thinking_end_token, + ) + + +def find_last_user_index(messages: List[Dict[str, Any]]) -> int: + """Find the index of the last user/developer message.""" + last_user_index = -1 + for idx in range(len(messages) - 1, -1, -1): + if messages[idx].get("role") in ["user", "developer"]: + last_user_index = idx + break + return last_user_index + + +# ============================================================ +# Message Rendering +# ============================================================ + + +def render_message( + index: int, + messages: List[Dict[str, Any]], + thinking_mode: str, + drop_thinking: bool = True, + reasoning_effort: Optional[str] = None, +) -> str: + """ + Render a single message at the given index into its encoded string form. + + This is the core function that converts each message in the conversation + into the DeepSeek-V4 format. + + Args: + index: Index of the message to render. + messages: Full list of messages in the conversation. + thinking_mode: Either "chat" or "thinking". + drop_thinking: Whether to drop reasoning content from earlier turns. + reasoning_effort: Optional reasoning effort level ("max", "high", or None). + + Returns: + Encoded string for this message. + """ + assert 0 <= index < len(messages) + assert thinking_mode in ["chat", "thinking"], f"Invalid thinking_mode `{thinking_mode}`" + + prompt = "" + msg = messages[index] + last_user_idx = find_last_user_index(messages) + + role = msg.get("role") + content = msg.get("content") + tools = msg.get("tools") + response_format = msg.get("response_format") + tool_calls = msg.get("tool_calls") + reasoning_content = msg.get("reasoning_content") + wo_eos = msg.get("wo_eos", False) + + if tools: + tools = tools_from_openai_format(tools) + if tool_calls: + tool_calls = tool_calls_from_openai_format(tool_calls) + + # Reasoning effort prefix (only at index 0 in thinking mode with max effort) + assert reasoning_effort in ["max", None, "high"], f"Invalid reasoning effort: {reasoning_effort}" + if index == 0 and thinking_mode == "thinking" and reasoning_effort == "max": + prompt += REASONING_EFFORT_MAX + + if role == "system": + prompt += system_msg_template.format(content=content or "") + if tools: + prompt += "\n\n" + render_tools(tools) + if response_format: + prompt += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + elif role == "developer": + assert content, f"Invalid message for role `{role}`: {msg}" + + content_developer = USER_SP_TOKEN + content_developer += content + + if tools: + content_developer += "\n\n" + render_tools(tools) + if response_format: + content_developer += "\n\n" + response_format_template.format(schema=to_json(response_format)) + + prompt += user_msg_template.format(content=content_developer) + + elif role == "user": + prompt += USER_SP_TOKEN + + # Handle content blocks (tool results mixed with text) + content_blocks = msg.get("content_blocks") + if content_blocks: + parts = [] + for block in content_blocks: + block_type = block.get("type") + if block_type == "text": + parts.append(block.get("text", "")) + elif block_type == "tool_result": + tool_content = block.get("content", "") + if isinstance(tool_content, list): + text_parts = [] + for b in tool_content: + if b.get("type") == "text": + text_parts.append(b.get("text", "")) + else: + text_parts.append(f"[Unsupported {b.get('type')}]") + tool_content = "\n\n".join(text_parts) + parts.append(tool_output_template.format(content=tool_content)) + else: + parts.append(f"[Unsupported {block_type}]") + prompt += "\n\n".join(parts) + else: + prompt += content or "" + + elif role == "latest_reminder": + prompt += LATEST_REMINDER_SP_TOKEN + latest_reminder_msg_template.format(content=content) + + elif role == "tool": + raise NotImplementedError( + "deepseek_v4 merges tool messages into user; please preprocess with merge_tool_messages()" + ) + + elif role == "assistant": + thinking_part = "" + tc_content = "" + + if tool_calls: + tc_list = [ + tool_call_template.format( + dsml_token=dsml_token, name=tc.get("name"), arguments=encode_arguments_to_dsml(tc) + ) + for tc in tool_calls + ] + tc_content += "\n\n" + tool_calls_template.format( + dsml_token=dsml_token, + tool_calls="\n".join(tc_list), + tc_block_name=tool_calls_block_name, + ) + + summary_content = content or "" + rc = reasoning_content or "" + + # Check if previous message has a task - if so, this is a task output (no thinking) + prev_has_task = index - 1 >= 0 and messages[index - 1].get("task") is not None + + if thinking_mode == "thinking" and not prev_has_task: + if not drop_thinking or index > last_user_idx: + thinking_part = thinking_template.format(reasoning_content=rc) + thinking_end_token + else: + thinking_part = "" + + if wo_eos: + prompt += assistant_msg_wo_eos_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + prompt += assistant_msg_template.format( + reasoning=thinking_part, + content=summary_content, + tool_calls=tc_content, + ) + else: + raise NotImplementedError(f"Unknown role: {role}") + + # Append transition tokens based on what follows + if index + 1 < len(messages) and messages[index + 1].get("role") not in ["assistant", "latest_reminder"]: + return prompt + + task = messages[index].get("task") + if task is not None: + # Task special token for internal classification tasks + assert task in VALID_TASKS, f"Invalid task: '{task}'. Valid tasks are: {list(VALID_TASKS)}" + task_sp_token = DS_TASK_SP_TOKENS[task] + + if task != "action": + # Non-action tasks: append task sp token directly after the message + prompt += task_sp_token + else: + # Action task: append Assistant + thinking token + action sp token + prompt += ASSISTANT_SP_TOKEN + prompt += thinking_end_token if thinking_mode != "thinking" else thinking_start_token + prompt += task_sp_token + + elif messages[index].get("role") in ["user", "developer"]: + # Normal generation: append Assistant + thinking token + prompt += ASSISTANT_SP_TOKEN + if not drop_thinking and thinking_mode == "thinking": + prompt += thinking_start_token + elif drop_thinking and thinking_mode == "thinking" and index >= last_user_idx: + prompt += thinking_start_token + else: + prompt += thinking_end_token + + return prompt + + +# ============================================================ +# Preprocessing +# ============================================================ + + +def merge_tool_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Merge tool messages into the preceding user message using content_blocks format. + + DeepSeek-V4 does not have a standalone "tool" role; instead, tool results + are encoded as blocks within user messages. + + This function converts a standard OpenAI-format conversation (with separate + "tool" role messages) into V4 format where tool results are merged into + user messages. + + Args: + messages: List of message dicts in OpenAI format. + + Returns: + Processed message list with tool messages merged into user messages. + """ + merged: List[Dict[str, Any]] = [] + + for msg in messages: + msg = copy.deepcopy(msg) + role = msg.get("role") + + if role == "tool": + # Convert tool message to a user message with tool_result block + tool_block = { + "type": "tool_result", + "tool_use_id": msg.get("tool_call_id", ""), + "content": msg.get("content", ""), + } + # Merge into previous message if it's already a user (merged tool) + if merged and merged[-1].get("role") == "user" and "content_blocks" in merged[-1]: + merged[-1]["content_blocks"].append(tool_block) + else: + merged.append( + { + "role": "user", + "content_blocks": [tool_block], + } + ) + elif role == "user": + text_block = {"type": "text", "text": msg.get("content", "")} + if ( + merged + and merged[-1].get("role") == "user" + and "content_blocks" in merged[-1] + and merged[-1].get("task") is None + ): + merged[-1]["content_blocks"].append(text_block) + else: + new_msg = { + "role": "user", + "content": msg.get("content", ""), + "content_blocks": [text_block], + } + # Preserve extra fields (task, wo_eos, mask, etc.) + for key in ("task", "wo_eos", "mask"): + if key in msg: + new_msg[key] = msg[key] + merged.append(new_msg) + else: + merged.append(msg) + + return merged + + +def sort_tool_results_by_call_order(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Sort tool_result blocks within user messages by the order of tool_calls + in the preceding assistant message. + + Args: + messages: Preprocessed message list (after merge_tool_messages). + + Returns: + Message list with sorted tool result blocks. + """ + last_tool_call_order: Dict[str, int] = {} + + for msg in messages: + role = msg.get("role") + if role == "assistant" and msg.get("tool_calls"): + last_tool_call_order = {} + for idx, tc in enumerate(msg["tool_calls"]): + tc_id = tc.get("id") or tc.get("function", {}).get("id", "") + if tc_id: + last_tool_call_order[tc_id] = idx + + elif role == "user" and msg.get("content_blocks"): + tool_blocks = [b for b in msg["content_blocks"] if b.get("type") == "tool_result"] + if len(tool_blocks) > 1 and last_tool_call_order: + sorted_blocks = sorted(tool_blocks, key=lambda b: last_tool_call_order.get(b.get("tool_use_id", ""), 0)) + sorted_idx = 0 + new_blocks = [] + for block in msg["content_blocks"]: + if block.get("type") == "tool_result": + new_blocks.append(sorted_blocks[sorted_idx]) + sorted_idx += 1 + else: + new_blocks.append(block) + msg["content_blocks"] = new_blocks + + return messages + + +# ============================================================ +# Main Encoding Function +# ============================================================ + + +def encode_messages( + messages: List[Dict[str, Any]], + thinking_mode: str, + context: Optional[List[Dict[str, Any]]] = None, + drop_thinking: bool = True, + add_default_bos_token: bool = True, + reasoning_effort: Optional[str] = None, +) -> str: + """ + Encode a list of messages into the DeepSeek-V4 prompt format. + + This is the main entry point for encoding conversations. It handles: + - BOS token insertion + - Thinking mode with optional reasoning content dropping + - Tool message merging into user messages + - Multi-turn conversation context + + Args: + messages: List of message dicts to encode. + thinking_mode: Either "chat" or "thinking". + context: Optional preceding context messages (already encoded prefix). + drop_thinking: If True, drop reasoning_content from earlier assistant turns + (only keep reasoning for messages after the last user message). + add_default_bos_token: Whether to prepend BOS token at conversation start. + reasoning_effort: Optional reasoning effort level ("max", "high", or None). + + Returns: + The encoded prompt string. + """ + context = context if context else [] + + # Preprocess: merge tool messages and sort tool results + messages = merge_tool_messages(messages) + messages = sort_tool_results_by_call_order(context + messages)[len(context) :] + if context: + context = merge_tool_messages(context) + context = sort_tool_results_by_call_order(context) + + full_messages = context + messages + + prompt = bos_token if add_default_bos_token and len(context) == 0 else "" + + # Resolve drop_thinking: if any message has tools defined, don't drop thinking + effective_drop_thinking = drop_thinking + if any(m.get("tools") for m in full_messages): + effective_drop_thinking = False + + if thinking_mode == "thinking" and effective_drop_thinking: + full_messages = _drop_thinking_messages(full_messages) + # After dropping, recalculate how many messages to render + # (context may have shrunk too) + num_to_render = len(full_messages) - len(_drop_thinking_messages(context)) + context_len = len(full_messages) - num_to_render + else: + num_to_render = len(messages) + context_len = len(context) + + for idx in range(num_to_render): + prompt += render_message( + idx + context_len, + full_messages, + thinking_mode=thinking_mode, + drop_thinking=effective_drop_thinking, + reasoning_effort=reasoning_effort, + ) + + return prompt + + +def _drop_thinking_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """ + Drop reasoning_content and non-essential messages before the last user message. + + Behavior: + - Messages with role in ["user", "system", "tool", "latest_reminder"] are always kept. + - Messages at or after the last user index are always kept. + - Assistant messages before the last user get reasoning_content removed. + - Developer messages before the last user are dropped entirely. + """ + last_user_idx = find_last_user_index(messages) + result = [] + keep_roles = {"user", "system", "tool", "latest_reminder", "direct_search_results"} + + for idx, msg in enumerate(messages): + role = msg.get("role") + if role in keep_roles or idx >= last_user_idx: + result.append(msg) + elif role == "assistant": + msg = copy.copy(msg) + msg.pop("reasoning_content", None) + result.append(msg) + # developer and other roles before last_user_idx are dropped + + return result + + +# ============================================================ +# Parsing (Decoding model output) +# ============================================================ + + +def _read_until_stop(index: int, text: str, stop: List[str]) -> Tuple[int, str, Optional[str]]: + """ + Read text from index until one of the stop strings is found. + + Returns: + Tuple of (new_index, content_before_stop, matched_stop_string_or_None). + """ + min_pos = len(text) + matched_stop = None + + for s in stop: + pos = text.find(s, index) + if pos != -1 and pos < min_pos: + min_pos = pos + matched_stop = s + + if matched_stop: + content = text[index:min_pos] + return min_pos + len(matched_stop), content, matched_stop + else: + content = text[index:] + return len(text), content, None + + +def parse_tool_calls(index: int, text: str) -> Tuple[int, Optional[str], List[Dict[str, str]]]: + """ + Parse DSML tool calls from text starting at the given index. + + Args: + index: Starting position in text. + text: The full text to parse. + + Returns: + Tuple of (new_index, last_stop_token, list_of_tool_call_dicts). + Each tool call dict has "name" and "arguments" keys. + """ + tool_calls: List[Dict[str, Any]] = [] + stop_token = None + tool_calls_end_token = f"" + + while index < len(text): + index, _, stop_token = _read_until_stop(index, text, [f"<{dsml_token}invoke", tool_calls_end_token]) + if _ != ">\n": + raise ValueError(f"Tool call format error: expected '>\\n' but got '{_}'") + + if stop_token == tool_calls_end_token: + break + + if stop_token is None: + raise ValueError("Missing special token in tool calls") + + index, tool_name_content, stop_token = _read_until_stop( + index, text, [f"<{dsml_token}parameter", f"\n$', tool_name_content, flags=re.DOTALL) + if len(p_tool_name) != 1: + raise ValueError(f"Tool name format error: '{tool_name_content}'") + tool_name = p_tool_name[0] + + tool_args: Dict[str, Tuple[str, str]] = {} + while stop_token == f"<{dsml_token}parameter": + index, param_content, stop_token = _read_until_stop(index, text, [f"/{dsml_token}parameter"]) + + param_kv = re.findall(r'^ name="(.*?)" string="(true|false)">(.*?)<$', param_content, flags=re.DOTALL) + if len(param_kv) != 1: + raise ValueError(f"Parameter format error: '{param_content}'") + param_name, string, param_value = param_kv[0] + + if param_name in tool_args: + raise ValueError(f"Duplicate parameter name: '{param_name}'") + tool_args[param_name] = (param_value, string) + + index, content, stop_token = _read_until_stop( + index, text, [f"<{dsml_token}parameter", f"\n": + raise ValueError(f"Parameter format error: expected '>\\n' but got '{content}'") + + tool_call = decode_dsml_to_arguments(tool_name=tool_name, tool_args=tool_args) + tool_calls.append(tool_call) + + return index, stop_token, tool_calls + + +def parse_message_from_completion_text(text: str, thinking_mode: str) -> Dict[str, Any]: + """ + Parse a model completion text into a structured assistant message. + + This function takes the raw text output from the model (a single assistant turn) + and extracts: + - reasoning_content (thinking block) + - content (summary/response) + - tool_calls (if any) + + NOTE: This function is designed to parse only correctly formatted strings and + will raise ValueError for malformed output. + + Args: + text: The raw completion text (including EOS token). + thinking_mode: Either "chat" or "thinking". + + Returns: + Dict with keys: "role", "content", "reasoning_content", "tool_calls". + tool_calls are in OpenAI format. + """ + summary_content, reasoning_content, tool_calls = "", "", [] + index, stop_token = 0, None + tool_calls_start_token = f"\n\n<{dsml_token}{tool_calls_block_name}" + + is_thinking = thinking_mode == "thinking" + is_tool_calling = False + + if is_thinking: + index, content_delta, stop_token = _read_until_stop(index, text, [thinking_end_token, tool_calls_start_token]) + reasoning_content = content_delta + assert stop_token == thinking_end_token, "Invalid thinking format: missing " + + index, content_delta, stop_token = _read_until_stop(index, text, [eos_token, tool_calls_start_token]) + summary_content = content_delta + if stop_token == tool_calls_start_token: + is_tool_calling = True + else: + assert stop_token == eos_token, "Invalid format: missing EOS token" + + if is_tool_calling: + index, stop_token, tool_calls = parse_tool_calls(index, text) + + index, tool_ends_text, stop_token = _read_until_stop(index, text, [eos_token]) + assert not tool_ends_text, "Unexpected content after tool calls" + + assert len(text) == index and stop_token in [eos_token, None], "Unexpected content at end" + + for sp_token in [bos_token, eos_token, thinking_start_token, thinking_end_token, dsml_token]: + assert ( + sp_token not in summary_content and sp_token not in reasoning_content + ), f"Unexpected special token '{sp_token}' in content" + + return { + "role": "assistant", + "content": summary_content, + "reasoning_content": reasoning_content, + "tool_calls": tool_calls_to_openai_format(tool_calls), + } diff --git a/lightllm/models/deepseek_v4/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"" + + 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)